mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-09-21 11:11:26 -03:00
Compare commits
323 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 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 | |||
| da071e8452 | |||
| a0bb6df2b8 | |||
| 6f5c444ec5 | |||
| 20f66a4fe1 | |||
| 879745da53 | |||
| 3afec0a0be | |||
| 06c270a6e1 | |||
| 87e93636dc | |||
| 074d1f2e51 | |||
| 40f922b0e8 | |||
| a7214b6cff | |||
| 8ca66e72eb | |||
| 90be5799e4 | |||
| 1a93b0eca2 | |||
| c2360a35ad | |||
| 030a32f8fa | |||
| 25e72b43ce | |||
| 41e9883daa | |||
| ae461ebc81 | |||
| 3ebf256c5d | |||
| 0905e2be6e | |||
| bd380bc1a1 | |||
| cb4fd3a0e6 | |||
| bbe0acac5c | |||
| 45e7c25308 | |||
| 86aa1d8059 | |||
| 74254756ef | |||
| 259e08e47c | |||
| 6647c45731 | |||
| b614a5c447 | |||
| b80830913c | |||
| e57e11897e | |||
| 8a16034135 | |||
| 7fc3b7e5be | |||
| b0c7a1baae | |||
| 6411d83d46 | |||
| 74a063b0e5 | |||
| 96376e5cce | |||
| e7c26bf722 | |||
| cef4129fc9 | |||
| 0a28500848 | |||
| fc3f3f3bdb | |||
| fa58297973 | |||
| 5d1a22fb8f | |||
| d2f50f26f1 | |||
| 4a6042d0b4 | |||
| 846206d958 | |||
| 0daf4924f0 | |||
| d38a3d091d | |||
| 94dd08646d | |||
| 658f88ca48 | |||
| f53352efb2 | |||
| 38809a9d1b | |||
| 395682509c | |||
| ef3e7d7bf4 | |||
| c85b6b64a1 | |||
| 34c87d4934 | |||
| 93472e5d67 | |||
| ae185ee714 | |||
| 795036275a | |||
| d43ab6e32f | |||
| 280181f92e | |||
| f8d98934ad | |||
| 303cca0d85 | |||
| c2f16784b3 | |||
| 5bc6d8286c | |||
| 3f8381ffee | |||
| 1ca99294c9 | |||
| 680f0a57f5 | |||
| 94e3f54571 | |||
| 5c2b2aedcc | |||
| ebc31fb963 | |||
| 9659df6ad9 | |||
| 04d131e9dc | |||
| 78fe6282c7 | |||
| 0c00ee22fc | |||
| 5fd4946b1f | |||
| f1d3ac0cdc | |||
| e2c45905f0 | |||
| b2c68e6a65 | |||
| eb0f6dd3b6 | |||
| 0bf87f9092 | |||
| 1da2433bb2 | |||
| 2d6cf545b9 | |||
| 6a259a14fa | |||
| 41e1fd1e1f | |||
| 95fb3c7fc9 | |||
| 8237e5f9ea | |||
| aa75986178 | |||
| b887922055 | |||
| 68fa0f29c7 | |||
| d9d362c9c9 | |||
| d0bc4be0dc | |||
| 420530f532 | |||
| 3001f0f0ef | |||
| b2a1307d23 | |||
| 64da845a58 | |||
| 27027c4497 | |||
| 86c85c08ec | |||
| 196c8ffc3e | |||
| cfc95ee02a | |||
| 479fa36997 | |||
| 3e1216e9bc | |||
| 007883b7d1 | |||
| dc9200a12c | |||
| d2f955266d | |||
| 8e724538bd | |||
| 6fcdeb799d | |||
| 97b9b1f62b | |||
| 4bf9a4b640 | |||
| c5088772e8 | |||
| 56acefbd6c | |||
| 5ab06c4aae | |||
| c11f4b5c68 | |||
| 86376284f4 | |||
| 2b8a2fc7d8 | |||
| f26e1b41c8 | |||
| c1671af99f | |||
| ac7707d0f6 | |||
| 381cd710a2 | |||
| ad0d18cb79 | |||
| 7980ee77d0 | |||
| 916b8bb327 | |||
| 87e3d4dea9 | |||
| 76a913f5e0 | |||
| d8c192e647 | |||
| c453437620 | |||
| 720fa6d909 | |||
| b4f71089f4 | |||
| 83e6657ead | |||
| 7ea6df4111 | |||
| d9ab92602a | |||
| 5ffadaed31 | |||
| 24f5f7df5d | |||
| daf01fb1d6 | |||
| 0f11b6def9 | |||
| 7df83f44b8 | |||
| 169fa7bed6 | |||
| 027b504fe8 | |||
| 186ef4da78 | |||
| dc674098e7 | |||
| 9087b4b07c | |||
| 8e45c22d7a | |||
| 191c4e03cd | |||
| ab4154c57d | |||
| 28e93d12ff | |||
| 75e63c758b | |||
| 823f71f269 | |||
| 042dd4088d | |||
| eaa791a9eb | |||
| 2228627ff4 | |||
| 4c647ad9c8 | |||
| 8ca3e6c33f | |||
| dd6bdbf297 | |||
| b47dde87e4 | |||
| 99e65cccd8 | |||
| 3bdacb8f46 | |||
| b4f9c224d3 | |||
| 5ec0399c81 | |||
| b464fdc333 | |||
| 53825500db | |||
| f2ac790752 | |||
| 0d8805cdee | |||
| 656e24ac9b | |||
| 6718b37403 | |||
| c9e5e784fc | |||
| f92f958682 | |||
| f63fab0676 | |||
| cfc4903c0c | |||
| a527a847fe | |||
| 91b0bf8933 | |||
| 66d1c96783 | |||
| 986128076e | |||
| 1de0a53241 | |||
| 0ec7eaf606 | |||
| d9fcb0e92b | |||
| f49b4ba4db | |||
| 84e708328b | |||
| 125bed3f09 | |||
| 077e70169d | |||
| e6dc169a05 | |||
| f34c02756d | |||
| 1e4c315481 | |||
| a8283a0d00 | |||
| 55896669fc | |||
| e341e0b9d2 | |||
| e6538c83bb | |||
| 92e1285ea5 | |||
| 2aabd1d90e | |||
| 7b8b778f83 | |||
| 7c8dc57d55 | |||
| fe95fae5f2 | |||
| ce8a95abf7 | |||
| c8e7e543d6 | |||
| a9dbb15ffa | |||
| cf64043f7d | |||
| ccaff92c18 | |||
| 585b5c922a | |||
| ea80c2224c | |||
| 8b0f56c1a6 | |||
| 8022d12f03 | |||
| 3939f7f91b | |||
| aebf2e37dd | |||
| f53f859a71 | |||
| d916375abe | |||
| 57983df4bd | |||
| c68d7559a0 | |||
| 9a8f5bf2d6 | |||
| a8d742b031 | |||
| c27e4d1bfc | |||
| d15a8aa9a2 | |||
| 74a7d12ca4 | |||
| 2f94a9773e | |||
| 37bdfa21ea | |||
| f0bf2728c9 | |||
| dc715aa273 | |||
| 7ee2361e87 | |||
| e04c22f83f | |||
| 681cc13e90 | |||
| 090e0297d4 | |||
| 6f71335be4 | |||
| 7f51812c1e | |||
| a9dc4d7b9d | |||
| 5d50ddb5d4 | |||
| f86198d234 | |||
| ffe65d983c | |||
| b0b5be913c | |||
| 01efcbc584 | |||
| 02c249917a | |||
| 419bbc90b2 | |||
| b0c4510fdb |
@@ -1,201 +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, including starting/restarting the server, using Chrome DevTools MCP to interact with the web UI at http://127.0.0.1:8188/loras, and verifying frontend-to-backend functionality. Covers workflow validation, UI interaction testing, and integration testing between the standalone Python backend and the browser frontend.
|
||||
---
|
||||
|
||||
# LoRa Manager E2E Testing
|
||||
|
||||
This skill provides workflows and utilities for end-to-end testing of LoRa Manager using Chrome DevTools MCP.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
- LoRa Manager project cloned and dependencies installed (`pip install -r requirements.txt`)
|
||||
- Chrome browser available for debugging
|
||||
- Chrome DevTools MCP connected
|
||||
|
||||
## Quick Start Workflow
|
||||
|
||||
### 1. Start LoRa Manager Standalone
|
||||
|
||||
```python
|
||||
# Use the provided script to start the server
|
||||
python .agents/skills/lora-manager-e2e/scripts/start_server.py --port 8188
|
||||
```
|
||||
|
||||
Or manually:
|
||||
```bash
|
||||
cd /home/miao/workspace/ComfyUI/custom_nodes/ComfyUI-Lora-Manager
|
||||
python standalone.py --port 8188
|
||||
```
|
||||
|
||||
Wait for server ready message before proceeding.
|
||||
|
||||
### 2. Open Chrome Debug Mode
|
||||
|
||||
```bash
|
||||
# Chrome with remote debugging on port 9222
|
||||
google-chrome --remote-debugging-port=9222 --user-data-dir=/tmp/chrome-lora-manager http://127.0.0.1:8188/loras
|
||||
```
|
||||
|
||||
### 3. Connect Chrome DevTools MCP
|
||||
|
||||
Ensure the MCP server is connected to Chrome at `http://localhost:9222`.
|
||||
|
||||
### 4. 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:8188/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
|
||||
python .agents/skills/lora-manager-e2e/scripts/start_server.py --port 8188 --restart
|
||||
|
||||
# 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)
|
||||
```
|
||||
|
||||
## Available Scripts
|
||||
|
||||
### scripts/start_server.py
|
||||
|
||||
Starts or restarts the LoRa Manager standalone server.
|
||||
|
||||
```bash
|
||||
python scripts/start_server.py [--port PORT] [--restart] [--wait]
|
||||
```
|
||||
|
||||
Options:
|
||||
- `--port`: Server port (default: 8188)
|
||||
- `--restart`: Kill existing server before starting
|
||||
- `--wait`: Wait for server to be ready before exiting
|
||||
|
||||
### scripts/wait_for_server.py
|
||||
|
||||
Polls 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
|
||||
2. Close browser pages (keep at least one open)
|
||||
3. Clear temporary data if needed
|
||||
@@ -1,324 +0,0 @@
|
||||
# Chrome DevTools MCP Cheatsheet for LoRa Manager
|
||||
|
||||
Quick reference for common MCP commands used in LoRa Manager E2E testing.
|
||||
|
||||
## Navigation
|
||||
|
||||
```python
|
||||
# Navigate to LoRA list page
|
||||
navigate_page(type="url", url="http://127.0.0.1:8188/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:8188/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:8188/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}"
|
||||
```
|
||||
@@ -1,272 +0,0 @@
|
||||
# LoRa Manager E2E Test Scenarios
|
||||
|
||||
This document provides detailed test scenarios for end-to-end validation of LoRa Manager features.
|
||||
|
||||
## 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:8188/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:8188/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:8188/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 --restart --wait`
|
||||
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,193 +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.
|
||||
"""
|
||||
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
|
||||
|
||||
def run_test():
|
||||
"""Run example E2E test flow."""
|
||||
|
||||
print("=" * 60)
|
||||
print("LoRa Manager E2E Test Example")
|
||||
print("=" * 60)
|
||||
|
||||
# Step 1: Start server
|
||||
print("\n[1/5] Starting LoRa Manager standalone server...")
|
||||
result = subprocess.run(
|
||||
[sys.executable, "start_server.py", "--port", "8188", "--wait", "--timeout", "30"],
|
||||
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("google-chrome --remote-debugging-port=9222 --user-data-dir=/tmp/chrome-lora-manager http://127.0.0.1:8188/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("""
|
||||
MCP Commands to execute:
|
||||
1. navigate_page(type="url", url="http://127.0.0.1:8188/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("""
|
||||
Scenario: Change setting and verify after restart
|
||||
|
||||
Steps:
|
||||
1. Navigate to settings page
|
||||
- navigate_page(type="url", url="http://127.0.0.1:8188/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", "--restart", "--wait"])
|
||||
|
||||
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:8188/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("""
|
||||
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:8188/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)
|
||||
@@ -1,169 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Start or restart LoRa Manager standalone server for E2E testing.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
import socket
|
||||
import signal
|
||||
import os
|
||||
|
||||
|
||||
def find_server_process(port: int) -> list[int]:
|
||||
"""Find PIDs of processes listening on the given port."""
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["lsof", "-ti", f":{port}"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=False
|
||||
)
|
||||
if result.returncode == 0 and result.stdout.strip():
|
||||
return [int(pid) for pid in result.stdout.strip().split("\n") if pid]
|
||||
except FileNotFoundError:
|
||||
# lsof not available, try netstat
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["netstat", "-tlnp"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=False
|
||||
)
|
||||
pids = []
|
||||
for line in result.stdout.split("\n"):
|
||||
if f":{port}" in line:
|
||||
parts = line.split()
|
||||
for part in parts:
|
||||
if "/" in part:
|
||||
try:
|
||||
pid = int(part.split("/")[0])
|
||||
pids.append(pid)
|
||||
except ValueError:
|
||||
pass
|
||||
return pids
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
return []
|
||||
|
||||
|
||||
def kill_server(port: int) -> None:
|
||||
"""Kill processes using the specified port."""
|
||||
pids = find_server_process(port)
|
||||
for pid in pids:
|
||||
try:
|
||||
os.kill(pid, signal.SIGTERM)
|
||||
print(f"Sent SIGTERM to process {pid}")
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
|
||||
# Wait for processes to terminate
|
||||
time.sleep(1)
|
||||
|
||||
# Force kill if still running
|
||||
pids = find_server_process(port)
|
||||
for pid in pids:
|
||||
try:
|
||||
os.kill(pid, signal.SIGKILL)
|
||||
print(f"Sent SIGKILL to process {pid}")
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
|
||||
|
||||
def is_server_ready(port: int, timeout: float = 0.5) -> bool:
|
||||
"""Check if server is accepting connections."""
|
||||
try:
|
||||
with socket.create_connection(("127.0.0.1", port), timeout=timeout):
|
||||
return True
|
||||
except (socket.timeout, ConnectionRefusedError, OSError):
|
||||
return False
|
||||
|
||||
|
||||
def wait_for_server(port: int, timeout: int = 30) -> bool:
|
||||
"""Wait for server to become ready."""
|
||||
start = time.time()
|
||||
while time.time() - start < timeout:
|
||||
if is_server_ready(port):
|
||||
return True
|
||||
time.sleep(0.5)
|
||||
return False
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Start LoRa Manager standalone server for E2E testing"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--port",
|
||||
type=int,
|
||||
default=8188,
|
||||
help="Server port (default: 8188)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--restart",
|
||||
action="store_true",
|
||||
help="Kill existing server before starting"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--wait",
|
||||
action="store_true",
|
||||
help="Wait for server to be ready before exiting"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--timeout",
|
||||
type=int,
|
||||
default=30,
|
||||
help="Timeout for waiting (default: 30)"
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Get project root (parent of .agents directory)
|
||||
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
skill_dir = os.path.dirname(script_dir)
|
||||
project_root = os.path.dirname(os.path.dirname(os.path.dirname(skill_dir)))
|
||||
|
||||
# Restart if requested
|
||||
if args.restart:
|
||||
print(f"Killing existing server on port {args.port}...")
|
||||
kill_server(args.port)
|
||||
time.sleep(1)
|
||||
|
||||
# Check if already running
|
||||
if is_server_ready(args.port):
|
||||
print(f"Server already running on port {args.port}")
|
||||
return 0
|
||||
|
||||
# Start server
|
||||
print(f"Starting LoRa Manager standalone server on port {args.port}...")
|
||||
cmd = [sys.executable, "standalone.py", "--port", str(args.port)]
|
||||
|
||||
# Start in background
|
||||
process = subprocess.Popen(
|
||||
cmd,
|
||||
cwd=project_root,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
start_new_session=True
|
||||
)
|
||||
|
||||
print(f"Server process started with PID {process.pid}")
|
||||
|
||||
# Wait for ready if requested
|
||||
if args.wait:
|
||||
print(f"Waiting for server to be ready (timeout: {args.timeout}s)...")
|
||||
if wait_for_server(args.port, args.timeout):
|
||||
print(f"Server ready at http://127.0.0.1:{args.port}/loras")
|
||||
return 0
|
||||
else:
|
||||
print(f"Timeout waiting for server")
|
||||
return 1
|
||||
|
||||
print(f"Server starting at http://127.0.0.1:{args.port}/loras")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -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.
|
||||
- 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.
|
||||
- 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
|
||||
```
|
||||
|
||||
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
|
||||
|
||||
- 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`.
|
||||
- Cache root: `<settings_dir>/cache`.
|
||||
- Canonical cache files:
|
||||
|
||||
@@ -14,6 +14,7 @@ from typing import Any
|
||||
|
||||
SECRET_PATTERN = re.compile(r"(key|token|secret|password|auth|credential)", re.IGNORECASE)
|
||||
APP_NAME = "ComfyUI-LoRA-Manager"
|
||||
SETTINGS_DIR_ENV = "LORA_MANAGER_SETTINGS_DIR"
|
||||
CACHE_SQLITE = {
|
||||
"model": ("model", "{library}.sqlite"),
|
||||
"recipe": ("recipe", "{library}.sqlite"),
|
||||
@@ -30,6 +31,15 @@ CACHE_JSON = {
|
||||
|
||||
def main() -> int:
|
||||
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.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.")
|
||||
|
||||
args = parser.parse_args()
|
||||
if args.settings_path:
|
||||
os.environ[SETTINGS_DIR_ENV] = args.settings_path
|
||||
context = build_context()
|
||||
|
||||
if args.command == "summary":
|
||||
@@ -78,6 +90,11 @@ def build_context() -> dict[str, Any]:
|
||||
|
||||
|
||||
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()
|
||||
portable = repo_root / "settings.json"
|
||||
if portable.exists():
|
||||
|
||||
@@ -25,6 +25,7 @@ model_cache/
|
||||
reasonix.toml
|
||||
.reasonix/
|
||||
.codegraph/
|
||||
.playwright-mcp/
|
||||
|
||||
# Vue widgets development cache (but keep build output)
|
||||
vue-widgets/node_modules/
|
||||
|
||||
@@ -0,0 +1,202 @@
|
||||
---
|
||||
slug: undo-delete-staging
|
||||
status: drafting
|
||||
intent: clear
|
||||
review_required: false
|
||||
pending-action: write .omo/plans/undo-delete-staging.md
|
||||
approach: "Option B: delayed physical deletion with Undo. Backend: same-volume rename to per-root staging dir (.lm-pending-delete/) [updated 2026-08: model staging moved to a SIBLING dir inside each deleted model's own folder — see 'Symlink fix (2026-08)' under Decisions] + manifest JSON (batch_id, expires_at, staged->original map) + purge (30s TTL timer + startup sweep + opportunistic) + undo-delete endpoint + settings toggle 'skip undo'. Small files (recipes: JSON+preview) copy to global staging under settings dir instead of rename. Frontend: extend toast system with action button + 30s countdown; delete flows (single model / recipe / bulk / duplicates) consume batch_id from delete response and show Undo toast; expired undo -> 'undo expired' toast. Plus confirm-modal friction (C-friction, NO type-to-confirm): delete button delay-activation 1.5s + modal shows file size 'will free X GB' + Cancel gets initial focus. i18n keys + sync_translation_keys.py."
|
||||
---
|
||||
|
||||
# Draft: undo-delete-staging
|
||||
|
||||
## Components (topology ledger)
|
||||
<!-- Lock the SHAPE before depth. One row per top-level component that can succeed or fail independently. -->
|
||||
<!-- id | outcome (one line) | status: active|deferred | evidence path -->
|
||||
- backend staging module (stage/purge/undo + manifest + per-volume dir resolution) | new module, active | pending exploration: model_lifecycle_service.py delete_model / delete_model_artifacts
|
||||
- delete endpoints return batch_id (model/recipe/bulk/duplicates) | active | pending exploration: handlers + response shapes
|
||||
- undo-delete HTTP endpoint + route registration | active | pending exploration: route registrar pattern
|
||||
- purge scheduling (30s timer + startup sweep + opportunistic) | active | pending exploration: app on_startup hooks
|
||||
- settings toggle "skip undo window" | active | pending exploration: settings service read pattern
|
||||
- frontend toast extension (action button + countdown) | active | pending exploration: showToast impl
|
||||
- frontend delete flows consume batch_id + Undo toast | active | pending exploration: call sites
|
||||
- confirm-modal friction (delay-activate + size display + cancel focus) | active | pending exploration: modal focus behavior
|
||||
- i18n keys + sync_translation_keys.py | active | known
|
||||
|
||||
## Open assumptions (announced defaults)
|
||||
<!-- Record any default you adopt instead of asking, so the user can veto it at the gate. -->
|
||||
<!-- assumption | adopted default | rationale | reversible? -->
|
||||
- Undo window TTL = 30s | 30s balances space-freeing intent vs accident recovery | yes (constant)
|
||||
- Staging dir name: `.lm-pending-delete/` under each model root; recipes: `{settings_dir}/.lm-pending-delete/` | hidden, same-volume [updated 2026-08: same-volume is now guaranteed by sibling staging inside the model's own folder, not by the root location], consistent | yes
|
||||
- Staging failure falls back to existing hard delete | user intent is delete; staging is best-effort; hard delete likely fails identically under same conditions | yes
|
||||
- Purge on startup uses expires_at (not purge-all) so a <30s restart with live tab can still undo | robust, matches client-side timer | yes
|
||||
- Settings toggle label: "Delete permanently immediately (skip undo window)" | power users freeing space | yes
|
||||
- C-friction: delete button enabled after 1.5s + modal shows freed size; NO type-to-confirm (user vetoed) | user explicitly rejected type-to-confirm | n/a
|
||||
- Bulk/duplicates delete: one batch id for whole action, one undo restores all | simplest consistent semantics | yes
|
||||
|
||||
## Findings (cited - path:lines)
|
||||
|
||||
### Backend
|
||||
- `delete_model_artifacts` (py/services/model_lifecycle_service.py:19-48) = physical delete via os.remove; patterns: main file + `{name}.metadata.json` + PREVIEW_EXTENSIONS (py/utils/constants.py:22-37). ALSO called by ModelScanner.bulk_delete_models (py/services/model_scanner.py:2221) - single swap point covers bulk models.
|
||||
- `ModelLifecycleService.delete_model` (model_lifecycle_service.py:101-154): fetches `cached_entry` (111-116) - SNAPSHOT available for cache restore; after delete: cache.raw_data removal + resort + bump_cache_version (136-143), `_hash_index.remove_by_path` (145-146), `_sync_update_for_model` (148; update-service only, no recipe JSON rewrites - recipe refs are hash-based, re-resolve on restore), `_persist_current_cache` (150-152), returns `{"success": True, "deleted_files": [...]}` (154).
|
||||
- Handler `delete_model` (py/routes/handlers/model_handlers.py:478-492): POST /api/lm/{prefix}/delete; response passthrough; `_broadcast_models_changed()` (57-74) after success; 400 `{"success":false,"error"}`; 500 plain text.
|
||||
- Recipe delete: handler (recipe_handlers.py:1422-1438) DELETE /api/lm/recipe/{recipe_id} -> persistence_service.delete_recipe (py/services/recipes/persistence_service.py:193-209): os.remove(recipe_json_path) + os.remove(image_path) (204-206), recipe_scanner.remove_recipe (208), returns `{"success": true, "message": ...}`. PersistenceResult dataclass (20-25).
|
||||
- Bulk models: POST /api/lm/{prefix}/bulk-delete (model_route_registrar.py:39) -> handler (model_handlers.py:974-994) -> lifecycle_service.bulk_delete_models (model_lifecycle_service.py:308-318) -> scanner.bulk_delete_models (model_scanner.py:2181-2269) which calls delete_model_artifacts per file (2221) + `_batch_update_cache_for_deleted_models` (2271-2335); response `{"success","status","total_deleted","total_attempted","cache_updated","results"}` (2254-2269).
|
||||
- Bulk recipes: POST /api/lm/recipes/bulk-delete (recipe_route_registrar.py:50) -> handler (recipe_handlers.py:1554-1573) -> persistence_service.bulk_delete (persistence_service.py:439-482): per-id os.remove x2 (464-466), recipe_scanner.bulk_remove (472); response `{"success","deleted","failed","total_deleted","total_failed"}` (474-482).
|
||||
- Duplicates: NO dedicated delete endpoints (find-only: GET /api/lm/{prefix}/find-duplicates model_route_registrar.py:59, GET /api/lm/recipes/find-duplicates recipe_route_registrar.py:49). Duplicate deletion reuses bulk-delete endpoints.
|
||||
- Startup hooks: lora_manager.py:183-187 `app.on_startup.append(lambda app: cls._initialize_services())` (ComfyUI mode, app = PromptServer.instance.app at :78); standalone.py:370-374 same (StandaloneLoraManager.add_routes). Background tasks: `asyncio.create_task(name=...)` (lora_manager.py:224-239; recipe_handlers.py:793). Singleton+asyncio.Lock pattern: model_scanner.py:40-63.
|
||||
- Settings: DEFAULT_SETTINGS (py/services/settings_manager.py:57-119), `get(key, default)` (1390-1392), get_settings_manager() (2215-2228), reset_settings_manager() (2231). Typed-bool getter example: get_skip_previously_downloaded_model_versions (1253-1262). Handlers: base_model_routes.py:70, base_recipe_routes.py:54.
|
||||
- Model roots: ModelScanner.get_model_roots base NotImplementedError (model_scanner.py:1073-1075); impls lora_scanner.py:31-45, checkpoint_scanner.py:428-441, embedding_scanner.py:24-36. `_find_root_for_file(file_path)` (model_scanner.py:1108-1124) returns containing root - for per-root staging dir computation [updated 2026-08: staging no longer uses the containing root; batches are siblings inside the model's own folder]. Business-path rule (AGENTS.md): use os.path.abspath, never realpath, for staging/undo routing.
|
||||
- Cache restore methods: ModelCache has raw_data + resort (conftest mocks: tests/conftest.py:144-154); ModelHashIndex.add_entry(sha256, file_path, autov3) (py/services/model_hash_index.py:16); RecipeScanner.add_recipe(recipe_data) (recipe_scanner.py:2136) -> recipe_cache.add_recipe (recipe_cache.py:64). No single-file incremental model rescan - use snapshot restore instead of rescan.
|
||||
- Route registrar: model_route_registrar.py:177 add_route(method, path, handler), :180 add_prefixed_route - undo endpoint can be a non-prefixed route via add_route.
|
||||
- Tests: tests/services/test_model_lifecycle_service.py (inline tmp_path files, per-test stub scanners ScannerForDelete/VersionAwareScanner etc); conftest MockScanner/MockCache/MockHashIndex (tests/conftest.py:134-212); integration fixtures tests/integration/conftest.py; lifecycle hook tests tests/routes/test_lora_manager_lifecycle.py:177-178, tests/standalone/test_standalone_server.py:83-84.
|
||||
|
||||
### Frontend
|
||||
- 5 delete call sites:
|
||||
a) Single model: static/js/utils/modalUtils.js confirmDelete (27-42) -> getModelApiClient().deleteModel(path); ignores return.
|
||||
b) Recipe single: static/js/components/RecipeCard.js confirmDeleteRecipe (405-449) - RAW fetch DELETE /api/lm/recipe/{id}, checks only response.ok, showToast toast.recipes.deletedSuccessfully, state.virtualScroller.removeItemByFilePath.
|
||||
c) Bulk: static/js/managers/BulkManager.js confirmBulkDelete (633-672) -> getActiveApiClient() (134-142) -> bulkDeleteModels(filePaths); reads result.cancelled/success/deleted_count/error.
|
||||
d) Recipe duplicates: static/js/components/DuplicatesManager.js confirmDeleteDuplicates (457-494) - RAW fetch POST /api/lm/recipes/bulk-delete, reads data.success/data.total_deleted, exitDuplicateMode().
|
||||
e) Model duplicates: static/js/components/ModelDuplicatesManager.js confirmDeleteDuplicates (710-776) - RAW fetch POST /api/lm/{type}/bulk-delete, reads data.total_deleted, then resetAndReload(true) + find-duplicates re-check.
|
||||
Bonus: static/js/components/shared/ModelVersionsTab.js:1136-1144 client.deleteModel (ignores return).
|
||||
- API clients: BaseModelApiClient.deleteModel (static/js/api/baseModelApi.js:184-216) returns true/false, shows its own toasts, does removeItemByFilePath inside; bulkDeleteModels (1591-1642) returns {success, deleted_count, failed_count, errors} or {success:false, cancelled:true}; RecipeSidebarApiClient.bulkDeleteModels (recipeApi.js:623-664) returns {success, deleted_count: total_deleted, ...}. Endpoint map apiConfig.js:56,64.
|
||||
- Toast: showToast(key, params={}, type='info', fallback=null) (static/js/utils/uiHelpers.js:136-193) - textContent only, NO action/button support; durations 2000/5000ms; CSS static/css/components/toast.css (.toast flex gap:12px - button can be added). Closest action pattern: bannerService.registerBanner actions array + onRegister (static/js/managers/BannerService.js; used uiHelpers.js:18-57).
|
||||
- i18n: locales/en.json delete keys (1303-1314 bulkDelete, 1945-1948 recipes, 1987-1991 models, 2124-2130 duplicates, 2166-2170 toast.api); t()/interpolate (static/js/i18n/index.js:193-248); translate wrapper (utils/i18nHelpers.js:13-23); sync script scripts/sync_translation_keys.py (en reference, [TODO: Translate] placeholders).
|
||||
- Refresh after undo: recipes -> window.recipeManager.loadRecipes(true) (recipes.js:359; used by FilterManager.js:752 etc) or refreshRecipes (recipeApi.js:308); models -> resetAndReload(true) from modelApiFactory (used by ModelDuplicatesManager.js:740).
|
||||
- Size for modal: card.dataset.file_size (ModelCard.js:467), formatFileSize (ModelModal.js:615).
|
||||
- Tests: tests/frontend/utils/uiHelpers.dom.test.js (toast), api/recipeApi.bulk.test.js, components/duplicatesManager.test.js, components/modelDuplicatesManager.test.js, pages/*Page.test.js, i18n tests tests/i18n/test_i18n.py.
|
||||
|
||||
## Decisions (with rationale)
|
||||
|
||||
1. Same-volume rename staging for model files (atomic, no copy cost for multi-GB files); cross-volume rename forbidden. [CORRECTED 2026-08: "same-volume because under the containing root" was only true for plain directories — nested symlinked subdirs could cross volumes. Superseded by sibling staging: `.lm-pending-delete/<batch_id>/` inside the deleted model's own folder makes stage/undo same-device by construction; see "Symlink fix (2026-08)" below.]
|
||||
2. Copy-to-global-staging for recipes (small files; avoids recipe JSON vs preview image cross-volume problem).
|
||||
3. Manifest JSON files are the only state - no DB changes. Manifest includes model cached_entry snapshot for exact cache restore (no rescan needed).
|
||||
4. Undo endpoint returns restored paths; expired batch -> 404-style error -> frontend 'undo expired' toast.
|
||||
5. Skip-undo setting honored server-side (no batch_id in response -> no undo toast client-side).
|
||||
6. Staging failure falls back to existing hard delete (best-effort undo, never blocks delete).
|
||||
7. Undo window TTL = 30s constant (PENDING_DELETE_TTL_SECONDS); startup sweep uses expires_at (survives restart; browser-tab timer survives).
|
||||
8. Purge triple-trigger: per-batch asyncio timer task + on_startup sweep + opportunistic purge at each stage/undo.
|
||||
9. Frontend: new showActionToast (keep showToast signature untouched; extract shared createToastElement/appendToast internals); undo click -> shared handleUndoDelete(batchId, refreshFn); full list refresh after undo (recipes: window.recipeManager.loadRecipes(true); models: resetAndReload(true)).
|
||||
10. C-friction wave (NO type-to-confirm - user vetoed): delete buttons delay-activate 1.5s after modal open, initial focus on Cancel, model delete modal gains "permanently deleted from disk" warning + file size display (card.dataset.file_size + formatFileSize).
|
||||
11. Model cache restore on undo: append snapshot to cache.raw_data (dedupe by file_path) + resort + bump_cache_version + _persist_current_cache + _hash_index.add_entry + _broadcast_models_changed. Recipe restore: copy back files + recipe_scanner.add_recipe(recipe_data loaded from restored JSON).
|
||||
|
||||
### Symlink fix (2026-08)
|
||||
|
||||
Post-execution addendum (plan `.omo/plans/undo-delete-symlink-fix.md`, commits 5fd4946b / 0c00ee22):
|
||||
|
||||
12. Model staging moved from `<model_root>/.lm-pending-delete/<batch_id>/` to `<model_dir>/.lm-pending-delete/<batch_id>/` (sibling of the model artifacts, inside the deleted model's own folder). Stage/undo renames are same-device BY CONSTRUCTION — EXDEV is impossible even when the business path traverses nested symlinks to other volumes (the decision-1 "containing root" guarantee covered only plain directories). EXDEV remains possible only for cross-volume merges, which keep the batch_ids-array fallback. Accepted edge: deleting the model's whole FOLDER during the 30s window destroys that batch (undo returns 404). Batch discovery uses an in-memory registry (`_known_batch_dirs`) with a startup reconciliation scan (`purge_expired(scan_roots=True)`) covering restarts and crash leftovers. Recipe batches unchanged (copy-based settings-dir staging with the `_restore_file` EXDEV fallback).
|
||||
|
||||
## Scope IN
|
||||
|
||||
- Model single delete (model_handlers delete_model / model_lifecycle_service)
|
||||
- Recipe delete (recipe_handlers delete_recipe / persistence_service)
|
||||
- Bulk delete (models scanner + recipes persistence) + duplicates (reuse bulk endpoints)
|
||||
- Undo endpoint POST /api/lm/undo-delete (models + recipes, one batch space)
|
||||
- Purge: timer + startup sweep + opportunistic
|
||||
- Settings toggle delete_undo_enabled + settings page checkbox
|
||||
- Frontend: showActionToast + all 5 delete flows + shared undo handler
|
||||
- C-friction modal changes (delay-activate + cancel focus + warning copy + size display)
|
||||
- i18n keys + sync_translation_keys.py
|
||||
- Backend + frontend tests
|
||||
|
||||
## Scope OUT (Must NOT have)
|
||||
|
||||
- NO type-to-confirm / hold-to-confirm friction (user vetoed)
|
||||
- NO OS trash integration (send2trash) in this iteration
|
||||
- NO persistent recycle-bin UI (no trash browsing page)
|
||||
- NO changes to exclude/unexclude flow
|
||||
- NO DB migrations
|
||||
- NO new dependencies (no send2trash)
|
||||
- NO changes to download flows
|
||||
- NO recipe-JSON rewriting on model undo (hash-based refs re-resolve themselves)
|
||||
|
||||
## Open questions
|
||||
|
||||
None - all implementation details resolved by exploration. Design decisions settled in conversation (B+C, no type-to-confirm).
|
||||
|
||||
## Approval gate
|
||||
status: approved
|
||||
<!-- Approach approved -> rerun scaffold without --draft-only, run Metis gap analysis, APPEND todo batches, fill TL;DR last, run structural self-check, then Phase 4 handoff. -->
|
||||
|
||||
## Review round state (ulw-plan-review-round-state-contract)
|
||||
```json
|
||||
{
|
||||
"transition": "replace",
|
||||
"phase": "review_round_initialized",
|
||||
"applies_when": ["retry_after_plan_change"],
|
||||
"atomic": true,
|
||||
"review_required": true,
|
||||
"plan_path": ".omo/plans/undo-delete-staging.md",
|
||||
"plan_sha256": "8cf7c9be38a76d8ef1fb832aba043d6d7e82b60465b1bf28c6eafa7045117adc",
|
||||
"review_round_id": "rr-undo-del-20260811-006",
|
||||
"round_status": "active",
|
||||
"pending-action": "review .omo/plans/undo-delete-staging.md",
|
||||
"review": {
|
||||
"momus": { "status": "pending", "workspace_root": "/mnt/data/reinstall-backup-2026-04-12/data/workspace/ComfyUI/custom_nodes/ComfyUI-Lora-Manager", "runtime_home": null, "target": ".omo/plans/undo-delete-staging.md", "round_id": "rr-undo-del-20260811-006", "plan_sha256": "8cf7c9be38a76d8ef1fb832aba043d6d7e82b60465b1bf28c6eafa7045117adc", "launch_id": null, "session": null, "result": null },
|
||||
"independent": { "status": "pending", "workspace_root": "/mnt/data/reinstall-backup-2026-04-12/data/workspace/ComfyUI/custom_nodes/ComfyUI-Lora-Manager", "runtime_home": null, "target": ".omo/plans/undo-delete-staging.md", "round_id": "rr-undo-del-20260811-006", "plan_sha256": "8cf7c9be38a76d8ef1fb832aba043d6d7e82b60465b1bf28c6eafa7045117adc", "launch_id": null, "session": null, "result": null }
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Review results + fix/retry ledger
|
||||
|
||||
### Round 1 (rr-undo-del-20260811-001, plan sha256 6c52bf99...)
|
||||
- momus: APPROVE (non-blocking notes: todo1+7 duplicate DEFAULT_SETTINGS key -> fixed todo 7 to verify-only; "batch_ids" plural in todos 8/9 acceptance -> fixed; purge OSError note -> folded into todo 1 purge semantics)
|
||||
- independent (oracle): CHANGES_REQUESTED
|
||||
- BLOCKING S1: scanner walks would index .lm-pending-delete staged files as ghost entries -> fixed: todo 1 now mandates scanner walk exclusion at model_scanner.py:706/:867/:1404/_process_model_file + acceptance (o) scanner-visibility test
|
||||
- BLOCKING S2: manifest lacks model_type, undo could restore into wrong cache/hash index -> fixed: manifest now carries model_type + todo 5 resolves per-type scanner via registrar pattern + acceptance (b) checkpoint-batch test
|
||||
- S3 merged-batch expires_at re-anchor -> fixed: merge_batches re-anchors now+TTL in todo 1 + todo 3/4 assertions
|
||||
- S4 manifest-less dir policy -> fixed: quarantine to <batch_id>.orphaned, never delete (todo 1 + acceptance g)
|
||||
- S5 partial-undo retry semantics -> fixed: per-entry restored flag write-through + retry test (acceptance e)
|
||||
- S6 purge locked-file failure semantics -> fixed: skip file, keep batch, never rmtree past errors (todo 1 + acceptance i)
|
||||
- T8 undo-after-restart test -> fixed: todo 5 acceptance (f)
|
||||
- T9 recipe undo -> re-delete test -> fixed: todo 5 acceptance (h)
|
||||
- T7 rescan-stale-entry test -> fixed: todo 5 acceptance (g)
|
||||
- Route registration pinned to shared routes class per mode (NOT per-model-type registrar which registers 3x) -> fixed: todo 5 now creates py/routes/pending_delete_routes.py registered once in lora_manager.py:170-172 + standalone.py:356-358 + duplicate-route test (e)
|
||||
- Version-index staleness on single-delete undo -> fixed: todo 5 follows bulk cache-update pattern incl. rebuild_version_index (model_scanner.py:2324)
|
||||
- Cancelled-bulk batch_id frontend handling -> fixed: todo 9 shows action toast on cancelled+staged-subset
|
||||
- Single-instance assumption -> added to Scope OUT
|
||||
- Occupied-refusal loss UX -> accepted-intent documented in success criteria + modal copy
|
||||
|
||||
### Round 2 (rr-undo-del-20260811-002, plan sha256 f3d52235...)
|
||||
- momus: APPROVE (all 12 round-1 fixes verified present; zero dead references; non-blocking nits only)
|
||||
- independent (oracle): CHANGES_REQUESTED
|
||||
- BLOCK-1: merge_batches file-movement semantics unspecified (silent data-loss vector) -> fixed: todo 1 now specifies move-into-winner-dir + entry re-point + loser-dirs-removed-only-when-empty + abort-on-move-failure (all batches intact) + merge inside service lock + acceptance (k) file-survival assertions + acceptance (l) merge-failure abort test
|
||||
- BLOCK-2: same-file parallel edits within waves (todo 5 vs 6 on lora_manager.py; todo 8 vs 9 on baseModelApi.js) -> fixed: waves/matrix now serialize 5->6 and 8->9 with explicit reasons; matrix updated
|
||||
- Recommended: checkpoint_scanner.py:331 exclusion -> fixed (todo 1 + acceptance p); S5 pre-check skips restored:true entries -> fixed (todo 1); _tags_count restore on undo -> fixed (todo 5 + acceptance j); undo-blind flows documented (ModelVersionsTab + misc_handlers:2456) -> fixed (todo 8 note + Scope OUT); merge-failure no-merge fallback contract (batch_ids array) -> fixed (todos 3/4/9)
|
||||
|
||||
### Round 3 (rr-undo-del-20260811-003, plan sha256 8f2dfd46...)
|
||||
- momus: APPROVE (all round-2 fixes verified present + spot-checked refs; no new contradictions)
|
||||
- independent (oracle): CHANGES_REQUESTED
|
||||
- BLOCKING A: merged batches never timer-purged after re-anchor (winner's original timer no-ops at old expiry; no fresh timer for re-anchored expiry; idle server -> merged batch lingers, violating "30s purge" success criterion; affects EVERY bulk delete) -> fixed: todo 1 merge_batches now ARMS A FRESH PURGE TIMER for the winner with re-anchored expiry + acceptance (q) fresh-timer test + purge_expired must enumerate ALL scanner types' roots (explicit in todo 1)
|
||||
- BLOCKING B: dependency matrix contradicted same-file policy for todos 8/9<->11 (5 shared files) and 12<->11 -> fixed: todo 11 now "Blocked by: 8, 9 (same files...)"; todo 12 blocked by 11 (sync after 11); wave text updated (11, then 12 AFTER 11); "Can parallelize with" columns corrected
|
||||
- BLOCKING C: frontend batch_ids sequential-undo fallback has NO test + merge->undo loser-restore + merge->purge assertions missing -> fixed: todo 9 acceptance now tests the batch_ids fallback path; todo 1 acceptance now has (k2)/(k3)
|
||||
- Notes folded: sub-second toast-tail expiry race accepted; EXDEV fallback = NORMAL path for cross-volume bulks [annotated 2026-08: after the sibling-staging fix, EXDEV can only arise during cross-volume MERGES, never during single stage/undo renames]
|
||||
|
||||
### Round 4 (rr-undo-del-20260811-004, plan sha256 179e7ff7...)
|
||||
- momus: APPROVE (round-3 fixes verified; one non-blocking nit: todo 11 inline "Blocked by: —" stale -> fixed to "8, 9")
|
||||
- independent (oracle): CHANGES_REQUESTED
|
||||
- BLOCKING GAP-1 (NEW, introduced by round-3 fix): todo 8 handleUndoDelete always-refresh/always-toast contract contradicted todo 9's sequential loop "exactly ONE final refresh" -> fixed: handleUndoDelete(batchId, refreshFn, {showToast, refresh}) suppression options; todo 9 loop uses suppressed calls + one final refresh/toast; acceptance extended (loop failure mid-way -> stop + error toast + no final refresh; 404 body discrimination expired vs occupied)
|
||||
- BLOCKING GAP-2: no cross-type purge enumeration test -> fixed: todo 1 acceptance (r) purges expired batches across lora root + checkpoint root + recipe staging dir in one call
|
||||
- Non-blocking folded: GAP-3 404-copy discrimination -> fixed in todo 8 (d); GAP-4 merge partial-failure rollback direction (move back + restore manifests, extended (l) asserts sequential constituent undo still restores everything) -> fixed in todo 1; GAP-5 post-restart timer-loss residual gap documented -> fixed in todo 6; GAP-6 usage_stats.py:424 walk added to exclusion mandate + todo 5 acceptance (k) embeddings undo test
|
||||
|
||||
### Round 5 (rr-undo-del-20260811-005, plan sha256 dfaa39ea...)
|
||||
- momus: APPROVE (all round-4 fixes verified; no new contradictions)
|
||||
- independent (oracle): CHANGES_REQUESTED
|
||||
- BLOCK-1: lock-ordering deadlock ambiguity (asyncio.Lock not re-entrant: opportunistic purge_expired called while stage/undo hold the lock would deadlock on first use) -> fixed: todo 1 now has explicit LOCK HIERARCHY (lock acquired ONLY by stage/merge/undo/purge_batch; purge_expired is lock-free and must be called BEFORE lock acquisition); todo 6 (c) updated with the same rule + acceptance (u) lock-no-deadlock test
|
||||
- BLOCK-2: purge edge semantics unspecified -> fixed: purge_batch treats missing staged files (partially-restored batches) as already-purged (FileNotFoundError silent no-op); sweep skips `.orphaned`-suffixed dirs (quarantine is terminal); acceptance (s) partially-restored purge + (t) quarantine-terminal tests
|
||||
- Non-blocking folded: todo 2/3 test-file collision -> todo 3's bulk tests moved to tests/services/test_model_scanner.py; todo 9 (d) DuplicatesManager refreshFn stated explicitly (recipes loadRecipes / models resetAndReload); modal-copy + bulk-count trade-offs acknowledged in success criteria; acceptance (r) extended with embeddings root
|
||||
|
||||
### Round 6 (rr-undo-del-20260811-006, plan sha256 8cf7c9be...)
|
||||
- momus: APPROVE (all round-5 fixes verified; no new contradictions; references verified)
|
||||
- independent (oracle): APPROVE — no blocking issues; all round-5 items fixed with working, tested solutions; no new race/data-loss/consistency defects
|
||||
- Deferred optional improvements (non-blocking, recorded for executor awareness; plan file left untouched to preserve the approved digest):
|
||||
1. Tag-count asymmetry: single delete_model never decrements _tags_count (lifecycle 101-154), bulk does (scanner 2297-2303); undo re-increment is exact for bulk, over-counts for single until rescan (cosmetic, self-healing). Optional fix riding in todo 2: decrement tags in the single-delete path to mirror bulk.
|
||||
2. Todo 5 factual nit: ModelCache.resort() already rebuilds the version index — explicit rebuild in undo is belt-and-braces, no action needed.
|
||||
3. Todo 8 premise nit: ModelVersionsTab call ignores deleteModel's return entirely — nothing breaks, no adaptation needed.
|
||||
4. Todo 3's pytest command includes test_model_lifecycle_service.py which todo 2 edits in the same wave — run that file's tests after todo 2 lands.
|
||||
5. merge_batches with a missing/quarantined constituent id: any sane fallback (abort -> batch_ids, or skip missing) acceptable — files stay staged either way.
|
||||
|
||||
## Review lifecycle
|
||||
- rounds: 6 (rr-undo-del-20260811-001..006); final round both lanes APPROVE
|
||||
- final live-plan validation: sha256 = 8cf7c9be38a76d8ef1fb832aba043d6d7e82b60465b1bf28c6eafa7045117adc — MATCHES approved round-6 digest
|
||||
- status: APPROVED — ready for execution handoff ($start-work undo-delete-staging)
|
||||
File diff suppressed because one or more lines are too long
@@ -2,6 +2,10 @@
|
||||
|
||||
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
|
||||
|
||||
### Backend Development
|
||||
@@ -28,16 +32,21 @@ 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=xml:coverage/backend/coverage.xml \
|
||||
--cov-report=json:coverage/backend/coverage.json
|
||||
```
|
||||
|
||||
### Frontend Development (Standalone Web UI)
|
||||
### Frontend Development (LoRA Manager Web UI)
|
||||
|
||||
```bash
|
||||
# Install dependencies (root and Vue widgets)
|
||||
npm install
|
||||
cd vue-widgets && npm install && cd ..
|
||||
|
||||
npm test # Run all tests (JS + Vue)
|
||||
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
|
||||
```
|
||||
|
||||
@@ -54,105 +63,198 @@ npm run test:watch # Watch mode
|
||||
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
|
||||
- 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`
|
||||
- PEP 8 with 4-space indentation, type hints required
|
||||
|
||||
### Naming Conventions
|
||||
#### Naming Conventions
|
||||
|
||||
- Files: `snake_case.py`, Classes: `PascalCase`, Functions/vars: `snake_case`
|
||||
- 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`
|
||||
- `async def` for I/O, `@pytest.mark.asyncio` for async tests
|
||||
- Singleton with `asyncio.Lock`: see `ModelScanner.get_instance()`
|
||||
- Return `aiohttp.web.json_response` or `web.Response`
|
||||
|
||||
### Testing
|
||||
### JavaScript/TypeScript
|
||||
|
||||
- `pytest` with `--import-mode=importlib`
|
||||
- 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
|
||||
#### Imports & Modules
|
||||
|
||||
- ES modules: `import { app } from "../../scripts/app.js"` for ComfyUI
|
||||
- Vue: `import { ref, computed } from 'vue'`, type imports: `import type { Foo }`
|
||||
- Export named functions: `export function foo() {}`
|
||||
|
||||
### Naming & Formatting
|
||||
#### Naming & Formatting
|
||||
|
||||
- camelCase for functions/vars/props, PascalCase for classes
|
||||
- Constants: `UPPER_SNAKE_CASE`, Files: `snake_case.js` or `kebab-case.js`
|
||||
- 2-space indentation preferred (follow existing file conventions)
|
||||
- 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)`
|
||||
- Event handlers via `addEventListener` or widget callbacks
|
||||
- Shared utilities: `web/comfyui/utils.js`
|
||||
- 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
|
||||
- Guard restoration loops with flag: `let isRestoring = false`
|
||||
- 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
|
||||
|
||||
- `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)
|
||||
- Handlers in `py/routes/handlers/` are pure functions with deps as params
|
||||
|
||||
### Model Types & Routes
|
||||
|
||||
- `BaseModelService` base for LoRA, Checkpoint, Embedding
|
||||
- `ModelScanner` for file discovery, hash deduplication
|
||||
- `PersistentModelCache` (SQLite) for persistence
|
||||
- Route registrars: `ModelRouteRegistrar`, endpoints: `/loras/*`, `/checkpoints/*`, `/embeddings/*`
|
||||
- WebSocket via `WebSocketManager` for real-time updates
|
||||
- API endpoints follow `/loras/*`, `/checkpoints/*`, `/embeddings/*` patterns
|
||||
- Route registrars organize endpoints by domain: `ModelRouteRegistrar`, `RecipeRouteRegistrar`, etc.
|
||||
- Request handlers in `py/routes/handlers/` implement route logic
|
||||
- All routes use aiohttp, return `web.json_response` or `web.Response`
|
||||
- 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
|
||||
|
||||
- Base: `py/recipes/base.py`, Enrichment: `RecipeEnrichmentService`
|
||||
- Parsers: `py/recipes/parsers/`
|
||||
- Base: `py/recipes/base.py`, Enrichment: `RecipeEnrichmentService` in `py/recipes/enrichment.py`
|
||||
- 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
|
||||
|
||||
### 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
|
||||
|
||||
- ALWAYS use English for comments (per copilot-instructions.md)
|
||||
- Dual mode: ComfyUI plugin (folder_paths) vs standalone (settings.json)
|
||||
- Detection: `os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"`
|
||||
- Run `python scripts/sync_translation_keys.py` after adding UI strings to `locales/en.json`
|
||||
- Symlinks require normalized paths
|
||||
|
||||
## 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. Standalone Web UI
|
||||
- Location: `./static/` and `./templates/`
|
||||
- Tech: Vanilla JS + CSS, served by standalone server
|
||||
- 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`
|
||||
- Symlinks require normalized paths.
|
||||
**Business paths vs real paths**: All stored paths and operation routing use the
|
||||
original paths as they appear under configured model roots — symlinks are NOT
|
||||
resolved. `os.path.realpath` is only for scanner dedup and the symlink cache.
|
||||
Any path passed to `os.remove`/`os.rename`/`shutil.move` or validated by a
|
||||
containment check MUST use the business path (i.e. `os.path.abspath`, not
|
||||
`realpath`).
|
||||
@@ -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
|
||||
+18
@@ -15,6 +15,10 @@ try: # pragma: no cover - import fallback for pytest collection
|
||||
from .py.nodes.lora_pool import LoraPoolLM
|
||||
from .py.nodes.lora_randomizer import LoraRandomizerLM
|
||||
from .py.nodes.lora_cycler import LoraCyclerLM
|
||||
from .py.nodes.lora_info import LoraInfoLM
|
||||
from .py.nodes.lora_syntax_to_path import LoraSyntaxToPath
|
||||
from .py.nodes.create_hook_lora import CreateHookLoraLM
|
||||
from .py.nodes.metadata_overwrite import MetadataOverwriteLM
|
||||
from .py.metadata_collector import init as init_metadata_collector
|
||||
except (
|
||||
ImportError
|
||||
@@ -56,6 +60,16 @@ except (
|
||||
"py.nodes.lora_randomizer"
|
||||
).LoraRandomizerLM
|
||||
LoraCyclerLM = importlib.import_module("py.nodes.lora_cycler").LoraCyclerLM
|
||||
LoraInfoLM = importlib.import_module("py.nodes.lora_info").LoraInfoLM
|
||||
LoraSyntaxToPath = importlib.import_module(
|
||||
"py.nodes.lora_syntax_to_path"
|
||||
).LoraSyntaxToPath
|
||||
CreateHookLoraLM = importlib.import_module(
|
||||
"py.nodes.create_hook_lora"
|
||||
).CreateHookLoraLM
|
||||
MetadataOverwriteLM = importlib.import_module(
|
||||
"py.nodes.metadata_overwrite"
|
||||
).MetadataOverwriteLM
|
||||
init_metadata_collector = importlib.import_module("py.metadata_collector").init
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
@@ -75,6 +89,10 @@ NODE_CLASS_MAPPINGS = {
|
||||
LoraPoolLM.NAME: LoraPoolLM,
|
||||
LoraRandomizerLM.NAME: LoraRandomizerLM,
|
||||
LoraCyclerLM.NAME: LoraCyclerLM,
|
||||
LoraInfoLM.NAME: LoraInfoLM,
|
||||
LoraSyntaxToPath.NAME: LoraSyntaxToPath,
|
||||
CreateHookLoraLM.NAME: CreateHookLoraLM,
|
||||
MetadataOverwriteLM.NAME: MetadataOverwriteLM,
|
||||
}
|
||||
|
||||
WEB_DIRECTORY = "./web/comfyui"
|
||||
|
||||
+520
-421
File diff suppressed because it is too large
Load Diff
@@ -54,7 +54,7 @@ The dedicated services encapsulate long-running work so handlers stay thin.
|
||||
| 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. |
|
||||
| `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`. |
|
||||
|
||||
## Maintaining critical invariants
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
# ComfyUI Dual-Mode Widget Rendering
|
||||
|
||||
ComfyUI custom node widgets render in one of two modes. Patterns that work in one often fail silently in the other. Test both.
|
||||
|
||||
## Mode Detection
|
||||
|
||||
```js
|
||||
typeof LiteGraph !== 'undefined' && LiteGraph.vueNodesMode
|
||||
```
|
||||
|
||||
In Vue SFCs, `window.LiteGraph` is unavailable — pass as a prop from `main.ts`.
|
||||
|
||||
## Canvas Mode Layout
|
||||
|
||||
Uses `computeLayoutSize()` + `distributeSpace()` to allocate widget height within the node. Widgets with `computeLayoutSize` participate in space distribution; those with `computeSize` have fixed height.
|
||||
|
||||
- `getMinHeight()` in `addDOMWidget` options → minimum widget height
|
||||
- `widget.computeLayoutSize()` → `{ minHeight, minWidth, maxHeight? }`
|
||||
- Avoid `getMaxHeight()` unless the widget genuinely needs a fixed cap (prevents user resize)
|
||||
|
||||
## Vue Mode Layout
|
||||
|
||||
Uses CSS Grid (`grid-template-rows`) + `ResizeObserver`. The ResizeObserver watches the widget's DOM and feeds back into grid row sizing. This creates a feedback loop: content grows → row resizes → more space for content → content reflows/grows → row resizes again.
|
||||
|
||||
### Height Containment
|
||||
|
||||
The fix: `contain: layout size` on the widget root. This tells the browser the element's intrinsic size is CSS-determined, not driven by descendant content. The ResizeObserver sees a stable size and the loop is broken.
|
||||
|
||||
```css
|
||||
.widget-root.lm-vue-node {
|
||||
height: 100%;
|
||||
min-height: var(--comfy-widget-min-height, 200px);
|
||||
contain: layout size;
|
||||
}
|
||||
```
|
||||
|
||||
Existing examples: `.lm-loras-container.lm-vue-node` and `.comfy-tags-container.lm-vue-node` in `web/comfyui/lm_styles.css`.
|
||||
|
||||
**Do NOT** fix height issues with `maxHeight`, `getMaxHeight()`, or inline `max-height` — these prevent the user from resizing the node.
|
||||
|
||||
## Scroll Wheel Isolation
|
||||
|
||||
Both modes need to distinguish "user wants to scroll widget content" from "user wants to zoom canvas".
|
||||
|
||||
**Canvas mode:** Add `@wheel` on widget root. Check `event.target.closest(selector)` for scrollable sub-areas. If scrollable → `event.stopPropagation()`. Otherwise → `app.canvas.processMouseWheel(event)`.
|
||||
|
||||
**Vue mode:** Add CSS class `lm-wheel-scrollable` to scrollable elements. The global capture-phase hook in `web/comfyui/utils.js` (`enableListWheelScroll`) detects wheel events on marked elements and manually scrolls them via `element.scrollTop`, consuming the event before canvas zoom sees it.
|
||||
|
||||
## DOM Structure
|
||||
|
||||
`main.ts` creates an outer `<div>` container, then `vueApp.mount(container)`. The Vue app renders its own root element inside.
|
||||
|
||||
- `container.id` / `container.style.*` → outer element
|
||||
- Vue scoped `<style>` → `[data-v-hash]` applies only to Vue root
|
||||
|
||||
Classes needed by scoped Vue CSS must go on the Vue root element. Pass data as props and bind with `:class` rather than manipulating the DOM from `main.ts`.
|
||||
|
||||
## Serialization
|
||||
|
||||
For stateful widgets that need workflow persistence:
|
||||
|
||||
- `serialize: true` in `addDOMWidget` options
|
||||
- `serializeValue()` → state snapshot (called on workflow save)
|
||||
- `onSetValue(v)` → restore state (called on workflow load)
|
||||
- Always handle missing keys in restored value for backward compatibility with old workflows
|
||||
@@ -0,0 +1,370 @@
|
||||
# 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, 1810 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.
|
||||
|
||||
---
|
||||
|
||||
## 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×你 → 您) |
|
||||
|
||||
---
|
||||
|
||||
## 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`), 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)
|
||||
@@ -39,6 +39,7 @@ These fields are present in all model metadata files.
|
||||
| `metadata_source` | string\|null | ❌ No | ✅ Yes | Last provider that supplied metadata (see below) |
|
||||
| `last_checked_at` | float | ❌ No (default: `0`) | ✅ Yes | Unix timestamp of last metadata check |
|
||||
| `hash_status` | string | ❌ No (default: `"completed"`) | ✅ Yes | Hash calculation status: `"pending"`, `"calculating"`, `"completed"`, `"failed"` |
|
||||
| `autov3` | string\|null | ❌ No | ✅ Yes | CivitAI AutoV3 hash (first 12 chars, lowercase hex) sourced from the safetensors embedded metadata (`sshs_model_hash` / `modelspec.hash_sha256`). **Absent** = not yet checked (may be backfilled later); **`null`** = checked but unavailable (header has no recognized hash); **12-char hex string** = value |
|
||||
|
||||
---
|
||||
|
||||
@@ -287,6 +288,7 @@ These fields are automatically synchronized with the filesystem:
|
||||
- `preview_url` — Updated if preview file is moved/removed
|
||||
- `sha256` — Updated during hash calculation (when `hash_status="pending"`)
|
||||
- `hash_status` — Updated during hash calculation
|
||||
- `autov3` — Set when metadata is first created (from safetensors header); may be backfilled later for entries where it is absent
|
||||
- `last_checked_at` — Timestamp of scan
|
||||
- `metadata_source` — Set based on metadata provider
|
||||
|
||||
@@ -345,6 +347,7 @@ These fields can be edited by users at any time through the Lora Manager UI or b
|
||||
| `metadata_source` | `null` |
|
||||
| `last_checked_at` | `0` |
|
||||
| `hash_status` | `"completed"` |
|
||||
| `autov3` | absent (not checked) or `null` (checked, no value) |
|
||||
| `usage_tips` | `"{}"` (LoRA only) |
|
||||
| `model_type` | `"checkpoint"` or `"embedding"` (not present in LoRA models) |
|
||||
|
||||
@@ -354,6 +357,7 @@ These fields can be edited by users at any time through the Lora Manager UI or b
|
||||
|
||||
| Version | Date | Changes |
|
||||
|---------|------|---------|
|
||||
| 1.1 | 2026-08 | Added `autov3` field (CivitAI AutoV3 hash with three-state semantics) |
|
||||
| 1.0 | 2026-03 | Initial schema documentation |
|
||||
|
||||
---
|
||||
|
||||
@@ -0,0 +1,206 @@
|
||||
# Plan: Multi-File Downloads Within a Single CivitAI Model Version
|
||||
|
||||
**Issue:** [#1058 — Cannot download multiple file variants from the same model version](https://github.com/willmiao/ComfyUI-Lora-Manager/issues/1058)
|
||||
**Status:** v2 — revised after adversarial review (backend correctness + frontend/tests)
|
||||
**Scope:** CivitAI/CivArchive downloads of `lora`, `checkpoint`, `embedding` model types. HuggingFace downloads are out of scope (already per-file).
|
||||
|
||||
> v2 changelog: incorporated 18 review findings. Key changes vs v1:
|
||||
> shared file resolver + `resolved_version_id` for the gate (R1); `file_params` normalization at API boundary (R2); D2 hash-matching rule fixed for empty-hash cases (R6/R7); D3 extended to re-point `version_index` on removal (R4); D4 replaced with a child table (R3); `delete_model_version` interaction documented (R5); `ModelVersionsTab` surface added to phase 2 (F6); phase-2 multi-file loop requires a reload-deferred download variant (F7); queue-retry `file_params=NULL` known issue recorded (R9); test-fixture gaps and revised estimates (F10).
|
||||
|
||||
---
|
||||
|
||||
## 1. Problem Statement
|
||||
|
||||
A CivitAI model version can contain multiple downloadable weight files (e.g. fp16/fp32, safetensors/ckpt, different sizes). LoRA Manager already has a working file-selection pipeline (frontend file dialog → `fileParams` → backend file matching), but downloaded state is tracked at the **model-version** level. After any single file of a version is downloaded:
|
||||
|
||||
1. The version is marked **In Library** and the file-selection entry point disappears.
|
||||
2. The backend rejects further download attempts for that version.
|
||||
|
||||
There is no way to download the remaining files of the same version through LoRA Manager.
|
||||
|
||||
## 2. Current State (verified against code; all references confirmed by review)
|
||||
|
||||
### 2.1 Download gating — backend (`py/services/download_manager.py`)
|
||||
|
||||
`_execute_original_download` enforces two version-level gates:
|
||||
|
||||
- **Library gate, early** (lines 1157–1184, before metadata fetch, fires when `model_version_id` given) and **late** (lines 1350–1376, fires only when `model_version_id is None`): `scanner.check_model_version_exists(version_id)` across lora/checkpoint/embedding scanners → hard error `"Model version already exists in ... library"`.
|
||||
- **History gate** (lines 1238–1279): when `skip_previously_downloaded_model_versions` setting is on, `_has_been_downloaded(model_type, version_id)` → silent skip. History DB primary key is `(model_type, version_id)` (`py/services/downloaded_version_history_service.py:61`).
|
||||
|
||||
File selection works: `file_params {id, type, format, size, fp}` is matched against `version_info.files` (lines 1498–1569), **but only under `if file_params and model_version_id:` (line 1499)** — with `model_id`-only requests the selection silently falls back to the primary file (1571–1619). `file_params` currently carries no file `name` or hash.
|
||||
|
||||
### 2.2 Downloaded-state surfacing — backend (`py/routes/handlers/model_handlers.py`)
|
||||
|
||||
`get_civitai_versions` (lines 2148–2188) sets per-version `existsLocally` via `cache.version_index.get(version_id)` (plus a single `localPath` from that entry) and `hasBeenDownloaded` via the history service. No per-file granularity.
|
||||
|
||||
### 2.3 Frontend blockers (`static/js/managers/DownloadManager.js`)
|
||||
|
||||
Three independent gates prevent re-entering the file dialog:
|
||||
|
||||
1. **Line 598:** file-select badge rendered only when `modelFiles.length > 1 && !existsLocally`.
|
||||
2. **Lines 666–681 (`updateNextButtonState`):** Next button disabled with "Already in Library" when `currentVersion.existsLocally`.
|
||||
3. **Lines 784–787 (`proceedToLocation`):** toast + abort when `currentVersion.existsLocally`.
|
||||
|
||||
The badge path (`confirmFileSelection` lines 737–759 → `proceedToLocationContent` → `startDownload` single mode → `executeDownloadWithProgress` → POST `file_params`, `static/js/api/baseModelApi.js:1236–1250`) has **zero** `existsLocally` guards (all 12 occurrences enumerated; none on this path; `import/DownloadManager.js` has none either). The `.exists-locally` CSS class is purely visual (`download-modal.css:496–499`). **Making the badge visible again is sufficient to unlock the flow** for phase 1.
|
||||
|
||||
Post-download refresh is clean: the modal closes and `resetAndReload(true)` performs a full library refetch (`DownloadManager.js:1063`); dialog reopen resets state and refetches versions with no client-side cache. No same-session staleness.
|
||||
|
||||
### 2.4 Local identity of the downloaded file
|
||||
|
||||
`LoraMetadata/CheckpointMetadata/EmbeddingMetadata.from_civitai_info(version_info, file_info, ...)` (`py/utils/models.py:245–369`) persists:
|
||||
|
||||
- `sha256` = `file_info.hashes.SHA256` (lowercased, defaults to `""`) — a stable per-file identity;
|
||||
- `civitai` = the full `version_info` payload (including the `files` list).
|
||||
|
||||
Metadata refresh (`metadata_sync_service.py:104–105`) replaces the `civitai` blob wholesale but never overwrites top-level `sha256`; `verify_duplicate_hashes` (481–526) corrects it to the on-disk hash. Top-level-sha256 matching is refresh-robust.
|
||||
|
||||
**Caveats (review R6/R7):**
|
||||
- SHA256 is not guaranteed: CivArchive's transform only sets `hashes` when source data carries it (`civarchive_client.py:185–189`); `from_civitai_info` defaults to `""`.
|
||||
- Name fallback is unreliable exactly when it matters: local `file_name` is extension-less (`models.py:264`) and `generate_unique_filename` rewrites it with a hash suffix on conflict (`download_manager.py:1125–1136`); checkpoints with `hash_status='pending'` keep empty sha256 until on-demand hashing (`model_scanner.py:1232–1240`).
|
||||
|
||||
### 2.5 Version index collision (pre-existing hazard)
|
||||
|
||||
`ModelCache.version_index` is single-valued (`model_cache.py:133`: `version_index[version_id] = item`). Two files of the same version in the library → second entry overwrites the first; `remove_from_version_index` (lines 151–181) drops the whole version key when the indexed entry is removed, even if a sibling file remains. ~10 read sites depend on this index (48 grep touch points total; readers include `recipe_scanner.py:2682–2726`, `recipe_format.py:37–40`, `misc_handlers.py:2440–2444`, `model_handlers.py`, `model_scanner.check_model_version_exists:2444`).
|
||||
|
||||
Review correction (F3): bulk paths `remove_models` (`model_scanner.py:2376`) and `update_single_model_cache` (`:1689`) call `rebuild_version_index()` right after, so a sibling re-enters the index in those flows — the hazard is narrower than v1 stated, but direct `remove_from_version_index` callers (e.g. `model_scanner.py:1018`) still drop the key, and the user-visible artifact in phase 1 is real: `localPath` in the dialog flips to whichever file was indexed last.
|
||||
|
||||
### 2.6 Entry points that send / don't send `file_params` (fully enumerated by review)
|
||||
|
||||
**Send `file_params` (user-initiated dialog flows only):** `DownloadManager.js:1611–1639` (single mode). API surface accepting arbitrary JSON `file_params`: GET `/api/lm/download-model-get` (`model_handlers.py:1634–1686`), POST `/api/lm/downloads/queue/add` (`model_handlers.py:1799–1832`).
|
||||
|
||||
**Never send `file_params` (keep version-level semantics):** batch download (`DownloadManager.js:1756–1766`; batch also filters out in-library versions at `:1648`), `downloadVersionWithDefaults` (`:1810–1830`), recipe import (`import/DownloadManager.js:269–276`), bulk missing-LoRA (`BulkMissingLoraDownloadManager.js:292–299`), `RecipeModal.js:1728–1736`, `ModelVersionsTab.js:1427`. `web/comfyui/` and `vue-widgets/src` contain **no** download triggers at all (grep-verified). `py/services/use_cases/` has only `download_model_use_case.py` (pass-through).
|
||||
|
||||
### 2.7 Paths that do NOT need changes (verified)
|
||||
|
||||
- **aria2 pause/resume** (`_resume_restored_aria2_download`, line 754+): resumes from persisted `resume_context`; never re-runs existence gates.
|
||||
- **`download_coordinator.py:90`**: pure pass-through of `file_params`.
|
||||
- **Update checker / plugin self-update** (`update_routes.py:496–501`): only closes the history DB handle.
|
||||
- **History delete semantics**: `mark_as_deleted` sets `is_deleted_override=1` and `has_been_downloaded` then returns False (`downloaded_version_history_service.py:276`) — LM-initiated deletes already reset the history skip.
|
||||
|
||||
### 2.8 Related pre-existing issues (record, not necessarily fix)
|
||||
|
||||
- **Queue retry drops file selection** (R9): `download_queue_service.retry_from_history` / `retry_all_failed` re-queue with `file_params=NULL` (`download_queue_service.py:705, 758`) although the queue table has a `file_params` column (`:43`) — a retried non-primary download silently reverts to the primary file. Fix alongside phase 1 (small: persist and reuse the column).
|
||||
- **`delete_model_version`** (`misc_handlers.py:2410–2487`): resolves the file via the single-valued `version_index` (2440–2444), deletes only that one file, and `mark_as_deleted` flags the **entire version** as deleted in history (2479) even when a sibling file remains in the library. See phase 2 item 6.1.5.
|
||||
|
||||
## 3. Goals / Non-Goals
|
||||
|
||||
**Goals**
|
||||
|
||||
- G1: A user can download any not-yet-downloaded file of a version already partially in the library (issue repro steps 6–8).
|
||||
- G2: True duplicates stay blocked: downloading the *same* file of the same version twice is rejected.
|
||||
- G3: Per-file downloaded state visible in the file dialog; multiple files selectable and downloadable in one pass.
|
||||
- G4: No regression for version-level semantics relied on by batch download, recipe missing-LoRA detection, and `skip_previously_downloaded_model_versions`.
|
||||
|
||||
**Non-Goals**
|
||||
|
||||
- No change to recipe `inLibrary` semantics ("any file of the version present" remains sufficient).
|
||||
- No change to the update-checker (version-level comparison).
|
||||
- No primary-key rebuild of the history database.
|
||||
- HuggingFace download flow untouched.
|
||||
|
||||
## 4. Design Decisions
|
||||
|
||||
- **D1 — Explicit file selection bypasses the history gate, version-level gates stay for everyone else.** The history skip exists to dedupe automated flows. A user explicitly picking a file is unambiguous intent; the file-level library gate (G2) still prevents real duplicates. **Guard conditions use normalized truthiness** (see D1a). All confirmed `file_params` senders are user-initiated dialog flows (2.6), and LM-initiated deletes already reset history (2.7), so the bypass only affects "downloaded but not LM-deleted" versions with the setting on — intended.
|
||||
- **D1a — `file_params` normalization at the boundary (R2).** `download-model-get` and `downloads/queue/add` accept arbitrary JSON; `{}` is `not None` but falsy and would bypass gates while downloading the primary file. Normalize `file_params = file_params or None` in the coordinator/handlers, and treat the bypass as active only when a target file id is resolvable.
|
||||
- **D2 — File identity matching rule (R6/R7):** hash-compare **only when both sides are non-empty** (lowercase SHA256 equality); name-compare when either side is empty. Never let `"" == ""` match. Name fallback caveats from 2.4 apply (renamed files, pending checkpoint hashes) — acceptable residual risk, worst case is a blocked re-download the user can retry after hashing completes.
|
||||
- **D3 — Cache indexes: additive multi-index + removal re-pointing (R4).** Add `version_files_index: Dict[int, List[dict]]` maintained alongside `version_index` by the same add/remove/rebuild methods; existing readers of `version_index` untouched. Additionally fix `remove_from_version_index`: when the popped entry has a surviving sibling (per the multi-index), re-point `version_index[version_id]` to the sibling instead of dropping the key; same for the `model_id_index` descriptor. This closes the 2.5 hazard for existing readers (`check_model_version_exists`, `existsLocally`, recipe matching) without restructuring anything.
|
||||
- **D4 — Per-file history via a child table (R3).** v1's additive-column approach is structurally impossible on a `(model_type, version_id)` PK (`ON CONFLICT DO UPDATE` would keep only the last file). Instead add `downloaded_version_files(model_type, version_id, file_id, file_name, downloaded_at, PRIMARY KEY(model_type, version_id, file_id))` — additive, no PK rebuild, honors the Non-Goal. Existing version-level table and queries unchanged. New per-file queries are opt-in. `_initialize_schema` uses `CREATE TABLE IF NOT EXISTS`, so the new table is created for existing DBs without any ALTER.
|
||||
- **D5 — UI flow reuse, with an extracted inner download function for multi-file (F7).** Phase 1 unlocks the existing badge → file dialog → location → download pipeline. Phase 2 upgrades the dialog to multi-select; iterating `executeDownloadWithProgress` as-is would produce N full library reloads, N toasts, and competing failure-summary modals — so phase 2 extracts a reload-deferred, failure-aggregating inner variant and runs one reload + one summary at the end.
|
||||
|
||||
## 5. Implementation — Phase 1 (fix the issue; independently shippable)
|
||||
|
||||
### 5.1 Backend — `py/services/download_manager.py`
|
||||
|
||||
1. **Normalize `file_params`** at the boundary (D1a): `download_coordinator.schedule_download` and the two API handlers (`model_handlers.py:1649–1666`, `1810–1832`) apply `file_params = file_params or None`.
|
||||
2. **Extract a shared file resolver** (R1): pull the matching logic at 1498–1569 into `_resolve_target_file(version_info, file_params) -> Optional[dict]`, used by **both** the new gate and the download-selection path. The selection path's condition (line 1499) switches from `model_version_id` to `resolved_version_id` (already computed at 1230–1236 from `version_info.id`), so gate and download always agree on the target file — including the `model_id`-only case.
|
||||
3. **New helper** `_find_local_file_entry(version_id, target_file) -> Optional[dict]`: iterate the three scanners' cached `raw_data` (NOT `version_index` — single-valued); candidates = entries whose `civitai.id` normalizes to `version_id`; match per D2.
|
||||
4. **Gate restructure in `_execute_original_download`**:
|
||||
- Early scanner gate (1157–1184): add `file_params is None` guard; with normalized `file_params`, defer (file identity not resolvable before metadata fetch).
|
||||
- After `version_info` fetch + `resolved_version_id` (~1229): when `file_params` present, resolve target file via the shared resolver; unresolvable → hard error "No matching file" (fail closed, prevents empty-dict bypass). Resolvable → `_find_local_file_entry`; hit → same hard error shape as today with the file name in the message.
|
||||
- History gate (1238–1279): add `file_params is None` (D1). Base-model skip (1281–1324) unchanged — still applies.
|
||||
- Late gate (1350–1376): add `file_params is None` guard (F2) — the post-fetch file-level check above already covers this case.
|
||||
- Nothing between the early gate and the post-fetch point assumes the version is absent (review task 6: only provider selection + metadata fetch; no DB writes; `_persist_aria2_state` runs only when actually downloading at 1659).
|
||||
5. **Queue retry fix** (2.8, small): persist `file_params` into the queue table on enqueue and reuse it in `retry_from_history` / `retry_all_failed`.
|
||||
6. Logging: `[download]` lines for file-level allow/block, consistent with existing style.
|
||||
|
||||
**Estimated:** ~150–220 LOC + resolver extraction.
|
||||
|
||||
### 5.2 Frontend — `static/js/managers/DownloadManager.js`
|
||||
|
||||
1. Line 598: drop `&& !existsLocally` from the badge condition (badge shows whenever `modelFiles.length > 1`).
|
||||
2. `fileParams` construction (1611–1616): add `name: this.selectedFile.name`.
|
||||
3. Surface the backend "file already in library" hard error as a toast instead of only the batch-summary modal (R10/F12 nit; reuse existing error message field).
|
||||
4. No changes to `updateNextButtonState` / `proceedToLocation` in phase 1; no template or CSS changes.
|
||||
|
||||
**Known phase-1 UX limitations (acknowledged, fixed in phase 2):** with all files downloaded the badge still renders and re-picking a downloaded file fails late (backend error after the location step); `localPath` may point at a sibling file; batch-preview "In Library" badge stays version-level and gives no hint of remaining files.
|
||||
|
||||
**Estimated:** ~10–30 LOC (confirmed realistic by review).
|
||||
|
||||
### 5.3 Phase 1 tests
|
||||
|
||||
Backend — extend `tests/services/test_download_manager_basic.py` (1694 lines; all fixture patterns exist):
|
||||
|
||||
- **Fixture gaps to add (F10):** `DummyScanner.get_cached_data()`/`raw_data` stub (~10 lines); `hashes.SHA256` in the metadata-provider payload's `files`.
|
||||
- Cases: same version + different SHA256 in library + `file_params` → proceeds; same SHA256 → hard error; `file_params=None` + version in library → hard error (unchanged); history-skip on + `file_params` → not skipped; without → skipped (unchanged); empty-dict `file_params` normalized → version-level behavior; `model_id`-only + `file_params` → gate and selection resolve the same file; legacy metadata (empty local sha256) matched by name; target file with empty SHA256 → name fallback, no `""==""` false positive.
|
||||
- Queue retry: `file_params` survives retry.
|
||||
- Assert proceed/abort via the existing `_execute_download` mock pattern.
|
||||
|
||||
Frontend (`tests/frontend/`): badge renders for multi-file version with `existsLocally=true` (pattern from `downloadManager.history.test.js`).
|
||||
|
||||
**Estimated:** ~150–250 LOC (confirmed realistic).
|
||||
|
||||
## 6. Implementation — Phase 2 (per-file status + multi-select + index hardening)
|
||||
|
||||
### 6.1 Backend
|
||||
|
||||
1. **`py/services/model_cache.py`** (D3): add `version_files_index`; maintain in `add_to_version_index` / `remove_from_version_index` / `rebuild_version_index`; removal re-points `version_index[version_id]` (and the `model_id_index` descriptor) to a surviving sibling instead of dropping the key.
|
||||
2. **`py/services/model_scanner.py`**: expose `get_files_for_version(version_id) -> List[dict]`.
|
||||
3. **`py/routes/handlers/model_handlers.py` `get_civitai_versions`**: annotate each version with `downloadedFiles: [{fileId, fileName, filePath}]` via `version_files_index` + D2 matching against `version.files`.
|
||||
4. **`py/services/downloaded_version_history_service.py`** (D4): new child table `downloaded_version_files`; `mark_downloaded` also upserts the child row when `file_id` known; `mark_as_deleted` clears the version's child rows only when no sibling remains in the library; new `get_downloaded_file_ids(model_type, version_id) -> set[int]`. `_record_downloaded_version_history` passes `file_info` through.
|
||||
5. **`delete_model_version`** (`misc_handlers.py:2410–2487`, R5): resolve **all** local files of the version via `version_files_index`; delete all (current endpoint semantics are version-level) or — if kept per-file — only `mark_as_deleted` when no sibling remains. Decide at implementation time; minimum is documenting current behavior.
|
||||
6. **`ModelVersionsTab` backend support**: none needed beyond item 3 (`downloadedFiles`); the tab consumes the same versions payload.
|
||||
|
||||
### 6.2 Frontend
|
||||
|
||||
1. **File dialog multi-select** — change surface (F8): option markup (`DownloadManager.js:712–724`), the single-select click handler (`727–734`), the `input[type="radio"]:checked` selector in `confirmFileSelection` (`738`); template `templates/components/modals/download_modal.html:48–60` (confirm-button label only); CSS `download-modal.css` — checkbox variant of `.file-option-radio input` (595–604) and a **new** `.file-option.disabled` style (does not exist). Files whose id ∈ `downloadedFiles` render disabled with an "In Library" tag.
|
||||
2. **Mixed-type guard (F8):** multi-select is restricted to files sharing the same routing target (`_isDiffusionModel` is computed once from a single `selectedFile` at 798–803; e.g. "Model" + "UNet" files route to different roots). Disallow mixed-type multi-select (simplest, predictable); single-file selection unchanged.
|
||||
3. **Multi-file download loop (D5/F7):** extract from `executeDownloadWithProgress` a reload-deferred, no-toast inner function; iterate per selected file with per-file progress; one `resetAndReload(true)` + one aggregated success/failure summary at the end (reuse `showDownloadBatchSummary`).
|
||||
4. **`updateNextButtonState` / `proceedToLocation`:** for multi-file versions, Next routes into the file dialog; hard block only when *every* weight file is downloaded.
|
||||
5. **`ModelVersionsTab.js` (F6):** the Download action (`:576` hidden when `isInLibrary`) — for multi-file versions with remaining files, show it and route into the download modal's file dialog; keep hidden when all files present.
|
||||
6. **Batch preview (F5):** `batch-preview-local-badge` (`:1320`) gains a "partially downloaded" hint for multi-file versions with remaining files.
|
||||
7. New i18n keys (`modals.download.fileSelection.inLibrary`, `downloadSelected`, partial-download tooltip, etc.) → run `python scripts/sync_translation_keys.py`.
|
||||
|
||||
### 6.3 Phase 2 tests
|
||||
|
||||
- `model_cache` (`tests/services/test_model_cache.py` already covers add/remove at 44–55): multi-valued index; sibling re-point on removal; rebuild.
|
||||
- `get_civitai_versions`: `downloadedFiles` correctness (hash match, name fallback, no match, CivArchive no-hash payload).
|
||||
- History service (`tests/services/test_downloaded_version_history_service.py` uses real SQLite on tmp_path): child-table creation on a legacy DB; per-file record/query; `mark_as_deleted` sibling semantics.
|
||||
- Frontend: dialog checkbox rendering/disabled state and multi-file confirm — **greenfield behavior coverage** (F10: no existing test exercises `showFileSelectionStep`/`confirmFileSelection`; infra exists, patterns must be built).
|
||||
|
||||
## 7. Risks and Mitigations
|
||||
|
||||
| Risk | Impact | Mitigation |
|
||||
|---|---|---|
|
||||
| History-gate bypass (D1) causes unwanted re-downloads in automated flows | Large checkpoint files re-downloaded | Bypass only with normalized, resolvable `file_params` (D1a); all such senders are user-initiated dialog flows (2.6, verified); tests pin batch/recipe/bulk behavior. |
|
||||
| Empty-hash matching edge cases (R6) | Duplicate download of the same file, or false block | D2 rule: hash only when both non-empty; name otherwise; never `""==""`. Residual risk documented (2.4). |
|
||||
| Phase-1 late-failure UX (F12) | User picks a downloaded file, fails only after location step | Toast surfacing (5.2.3); phase 2 disables downloaded files up front. |
|
||||
| Phase-2 index change corrupts existing behavior | Recipe matching, delete flows | Additive index + re-point only; `version_index` read semantics unchanged; `remove_models`/`update_single_model_cache` already rebuild (F3); tests. |
|
||||
| `delete_model_version` marks whole version deleted while sibling remains (R5) | History wrongly suppresses re-download of the surviving sibling's version | Phase 2 item 6.1.5; documented until then. |
|
||||
| History child-table migration failure on user installs | Service init crash | `CREATE TABLE IF NOT EXISTS` in `_initialize_schema`; failure degrades to version-level behavior (per-file queries return empty). |
|
||||
| Batch-preview badge misleading for partial versions (F5) | Minor UX confusion | Acknowledged in phase 1; fixed in phase 2 item 6.2.6. |
|
||||
| UI confusion: version shows "In Library" while files remain downloadable | Support burden | Phase 2: per-file disabled state + partial-download tooltip. |
|
||||
| Hash-identical sibling files (repacked content) | Second file blocked | Acceptable: scanner hash dedup already collapses them. |
|
||||
|
||||
## 8. Rollout
|
||||
|
||||
1. **Commit 1** — `fix(download): allow downloading additional files of an in-library model version (#1058)` → Phase 1 (5.1–5.3).
|
||||
2. **Commit 2** — `feat(download): per-file download status and multi-file selection (#1058)` → Phase 2 (6.1–6.3).
|
||||
|
||||
Phase 1 alone resolves the issue as reported; phase 2 can ship in a later release if review prefers smaller increments.
|
||||
|
||||
## 9. Effort Estimate (revised after review)
|
||||
|
||||
| Phase | Backend | Frontend | Tests | Risk |
|
||||
|---|---|---|---|---|
|
||||
| 1 | ~150–220 LOC (+ queue-retry fix ~30) | ~10–30 LOC | ~150–250 LOC | Low |
|
||||
| 2 | ~250–350 LOC | ~250–350 LOC (multi-file loop refactor + ModelVersionsTab + batch badge) | ~250–350 LOC (dialog tests greenfield) | Medium |
|
||||
@@ -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,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.
|
||||
File diff suppressed because one or more lines are too long
+568
-226
File diff suppressed because it is too large
Load Diff
+436
-94
File diff suppressed because it is too large
Load Diff
+577
-235
File diff suppressed because it is too large
Load Diff
+569
-227
File diff suppressed because it is too large
Load Diff
+596
-254
File diff suppressed because it is too large
Load Diff
+534
-192
File diff suppressed because it is too large
Load Diff
+538
-196
File diff suppressed because it is too large
Load Diff
+553
-211
File diff suppressed because it is too large
Load Diff
+464
-122
File diff suppressed because it is too large
Load Diff
+477
-135
File diff suppressed because it is too large
Load Diff
+62
-14
@@ -1,9 +1,13 @@
|
||||
# pyright: reportImportCycles=false
|
||||
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||
# import cycles. Breaking them would require an architectural refactor.
|
||||
import os
|
||||
import platform
|
||||
import posixpath
|
||||
import threading
|
||||
from pathlib import Path
|
||||
import folder_paths # type: ignore
|
||||
import folder_paths # pyright: ignore[reportMissingImports]
|
||||
from typing import Any, Dict, Iterable, List, Mapping, Optional, Set, Tuple
|
||||
import logging
|
||||
import json
|
||||
@@ -90,7 +94,7 @@ def _resolve_valid_default_root(
|
||||
|
||||
|
||||
def _normalize_folder_paths_for_comparison(
|
||||
folder_paths: Mapping[str, Iterable[str]],
|
||||
folder_paths: Mapping[str, Any],
|
||||
) -> Dict[str, Set[str]]:
|
||||
"""Normalize folder paths for comparison across libraries."""
|
||||
|
||||
@@ -208,6 +212,12 @@ class Config:
|
||||
if not isinstance(library_config, dict):
|
||||
return
|
||||
|
||||
# Always read recipes_path — it is independent of extra folder paths
|
||||
# and must be set before any early returns below.
|
||||
recipes_path = library_config.get("recipes_path", "")
|
||||
if isinstance(recipes_path, str) and recipes_path:
|
||||
self.recipes_path = recipes_path
|
||||
|
||||
extra_folder_paths = library_config.get("extra_folder_paths")
|
||||
if not isinstance(extra_folder_paths, dict):
|
||||
return
|
||||
@@ -233,10 +243,6 @@ class Config:
|
||||
extra_embedding
|
||||
)
|
||||
|
||||
recipes_path = library_config.get("recipes_path", "")
|
||||
if isinstance(recipes_path, str) and recipes_path:
|
||||
self.recipes_path = recipes_path
|
||||
|
||||
if self.extra_loras_roots:
|
||||
logger.info(
|
||||
"Found extra LoRA roots:"
|
||||
@@ -357,6 +363,47 @@ class Config:
|
||||
"Failed to rename legacy 'default' library: %s", rename_error
|
||||
)
|
||||
|
||||
# Clean up a stale "default" library entry that has no meaningful
|
||||
# paths configured (e.g. leftover bootstrap artifact). This only
|
||||
# fires when "comfyui" already exists so we never delete the last
|
||||
# remaining library.
|
||||
if (
|
||||
"default" in libraries
|
||||
and "comfyui" in libraries
|
||||
and isinstance(default_library, Mapping)
|
||||
):
|
||||
default_folder_paths = _normalize_library_folder_paths(
|
||||
default_library
|
||||
)
|
||||
default_extra_paths = default_library.get("extra_folder_paths", {})
|
||||
has_meaningful_paths = bool(default_folder_paths) or bool(
|
||||
default_extra_paths
|
||||
) or any(
|
||||
default_library.get(key)
|
||||
for key in (
|
||||
"default_lora_root",
|
||||
"default_checkpoint_root",
|
||||
"default_unet_root",
|
||||
"default_embedding_root",
|
||||
"recipes_path",
|
||||
)
|
||||
)
|
||||
if not has_meaningful_paths:
|
||||
try:
|
||||
settings_service.delete_library("default")
|
||||
libraries_changed = True
|
||||
logger.info(
|
||||
"Removed stale 'default' library entry "
|
||||
"with no meaningful paths configured"
|
||||
)
|
||||
libraries = settings_service.get_libraries()
|
||||
comfy_library = libraries.get("comfyui", {})
|
||||
except Exception as delete_error:
|
||||
logger.debug(
|
||||
"Failed to remove stale 'default' library: %s",
|
||||
delete_error,
|
||||
)
|
||||
|
||||
default_lora_root = _resolve_valid_default_root(
|
||||
comfy_library.get("default_lora_root", ""),
|
||||
list(self.loras_roots or []),
|
||||
@@ -439,7 +486,7 @@ class Config:
|
||||
import ctypes
|
||||
|
||||
FILE_ATTRIBUTE_REPARSE_POINT = 0x400
|
||||
attrs = ctypes.windll.kernel32.GetFileAttributesW(str(path)) # type: ignore[attr-defined]
|
||||
attrs = ctypes.windll.kernel32.GetFileAttributesW(str(path)) # pyright: ignore[reportAttributeAccessIssue]
|
||||
return attrs != -1 and (attrs & FILE_ATTRIBUTE_REPARSE_POINT)
|
||||
except Exception as e:
|
||||
logger.error(f"Error checking Windows reparse point: {e}")
|
||||
@@ -448,7 +495,7 @@ class Config:
|
||||
logger.error(f"Error checking link status for {path}: {e}")
|
||||
return False
|
||||
|
||||
def _entry_is_symlink(self, entry: os.DirEntry) -> bool:
|
||||
def _entry_is_symlink(self, entry: os.DirEntry[str]) -> bool:
|
||||
"""Check if a directory entry is a symlink, including Windows junctions."""
|
||||
if entry.is_symlink():
|
||||
return True
|
||||
@@ -457,7 +504,7 @@ class Config:
|
||||
import ctypes
|
||||
|
||||
FILE_ATTRIBUTE_REPARSE_POINT = 0x400
|
||||
attrs = ctypes.windll.kernel32.GetFileAttributesW(entry.path) # type: ignore[attr-defined]
|
||||
attrs = ctypes.windll.kernel32.GetFileAttributesW(entry.path) # pyright: ignore[reportAttributeAccessIssue]
|
||||
return attrs != -1 and (attrs & FILE_ATTRIBUTE_REPARSE_POINT)
|
||||
except Exception:
|
||||
pass
|
||||
@@ -1083,8 +1130,8 @@ class Config:
|
||||
|
||||
def _apply_library_paths(
|
||||
self,
|
||||
folder_paths: Mapping[str, Iterable[str]],
|
||||
extra_folder_paths: Optional[Mapping[str, Iterable[str]]] = None,
|
||||
folder_paths: Mapping[str, Any],
|
||||
extra_folder_paths: Optional[Mapping[str, Any]] = None,
|
||||
recipes_path: str = "",
|
||||
) -> None:
|
||||
self._path_mappings.clear()
|
||||
@@ -1389,12 +1436,13 @@ class Config:
|
||||
# ('_lm_config_cache') that is NEVER removed from sys.modules (its key does
|
||||
# NOT start with 'py.'), so it survives re-imports of py.* modules.
|
||||
_CONFIG_SENTINEL = "_lm_config_cache"
|
||||
config: Config
|
||||
if _CONFIG_SENTINEL in _sys.modules:
|
||||
# Re-import: reuse the existing singleton from the sentinel.
|
||||
config: Config = _sys.modules[_CONFIG_SENTINEL].config # type: ignore[valid-type]
|
||||
config = _sys.modules[_CONFIG_SENTINEL].config
|
||||
else:
|
||||
config: Config = Config()
|
||||
config = Config()
|
||||
# Register the sentinel so re-imports of py.config find us.
|
||||
_sentinel_mod = _types.ModuleType(_CONFIG_SENTINEL)
|
||||
_sentinel_mod.config = config
|
||||
setattr(_sentinel_mod, "config", config)
|
||||
_sys.modules[_CONFIG_SENTINEL] = _sentinel_mod
|
||||
|
||||
+18
-1
@@ -14,7 +14,7 @@ standalone_mode = (
|
||||
if not standalone_mode:
|
||||
setup_logging()
|
||||
|
||||
from server import PromptServer # type: ignore
|
||||
from server import PromptServer # pyright: ignore[reportMissingImports]
|
||||
|
||||
from .config import config
|
||||
from .services.model_service_factory import (
|
||||
@@ -25,10 +25,12 @@ from .routes.recipe_routes import RecipeRoutes
|
||||
from .routes.stats_routes import StatsRoutes
|
||||
from .routes.update_routes import UpdateRoutes
|
||||
from .routes.misc_routes import MiscRoutes
|
||||
from .routes.pending_delete_routes import PendingDeleteRoutes
|
||||
from .routes.preview_routes import PreviewRoutes
|
||||
from .routes.example_images_routes import ExampleImagesRoutes
|
||||
from .services.service_registry import ServiceRegistry
|
||||
from .services.settings_manager import get_settings_manager
|
||||
from .services.pending_delete_service import get_pending_delete_service
|
||||
from .utils.example_images_migration import ExampleImagesMigration
|
||||
from .services.websocket_manager import ws_manager
|
||||
from .services.example_images_cleanup_service import ExampleImagesCleanupService
|
||||
@@ -170,6 +172,7 @@ class LoraManager:
|
||||
RecipeRoutes.setup_routes(app)
|
||||
UpdateRoutes.setup_routes(app)
|
||||
MiscRoutes.setup_routes(app)
|
||||
PendingDeleteRoutes.setup_routes(app)
|
||||
ExampleImagesRoutes.setup_routes(app, ws_manager=ws_manager)
|
||||
PreviewRoutes.setup_routes(app)
|
||||
|
||||
@@ -245,6 +248,20 @@ class LoraManager:
|
||||
cls._run_post_initialization_tasks(init_tasks), name="post_init_tasks"
|
||||
)
|
||||
|
||||
# Startup sweep: purge pending-delete batches that expired during a
|
||||
# previous run. Non-blocking (fire-and-forget); purge_expired only
|
||||
# removes already-expired batches, so a staged undo that survived a
|
||||
# restart stays restorable. scan_roots=True runs the reconciliation
|
||||
# pass first so leftover batches (the in-process registry is empty
|
||||
# after a restart) are re-discovered on disk. Covers both plugin
|
||||
# and standalone modes (StandaloneLoraManager reuses this
|
||||
# classmethod).
|
||||
pending_delete_service = await get_pending_delete_service()
|
||||
asyncio.create_task(
|
||||
pending_delete_service.purge_expired(scan_roots=True),
|
||||
name="pending_delete_startup_sweep",
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
"LoRA Manager: All services initialized and background tasks scheduled"
|
||||
)
|
||||
|
||||
@@ -22,7 +22,7 @@ if not standalone_mode:
|
||||
|
||||
logger.info("ComfyUI Metadata Collector initialized")
|
||||
|
||||
def get_metadata(prompt_id=None): # type: ignore[no-redef]
|
||||
def get_metadata(prompt_id=None): # pyright: ignore[reportRedeclaration]
|
||||
"""Helper function to get metadata from the registry"""
|
||||
registry = MetadataRegistry()
|
||||
return registry.get_metadata(prompt_id)
|
||||
@@ -31,6 +31,6 @@ else:
|
||||
def init():
|
||||
logger.info("ComfyUI Metadata Collector disabled in standalone mode")
|
||||
|
||||
def get_metadata(prompt_id=None): # type: ignore[no-redef]
|
||||
def get_metadata(prompt_id=None): # pyright: ignore[reportRedeclaration]
|
||||
"""Dummy implementation for standalone mode"""
|
||||
return {}
|
||||
|
||||
@@ -1,5 +1,11 @@
|
||||
"""Constants used by the metadata collector"""
|
||||
|
||||
# Sentinel value for clip_skip to distinguish "unconnected / widget default"
|
||||
# from "user wired value 0". Both ComfyUI CLIPSetLastLayer (-24..-1) and
|
||||
# A1111 conventions treat 0 as meaningless for clip skipping, but users may
|
||||
# explicitly wire 0 to the overwrite node to express "no clip skip / default".
|
||||
CLIP_SKIP_SENTINEL = -25
|
||||
|
||||
# Metadata categories
|
||||
MODELS = "models"
|
||||
PROMPTS = "prompts"
|
||||
@@ -9,6 +15,14 @@ EMBEDDINGS = "embeddings"
|
||||
SIZE = "size"
|
||||
IMAGES = "images"
|
||||
IS_SAMPLER = "is_sampler" # New constant to mark sampler nodes
|
||||
OVERWRITE = "overwrite" # Manual metadata overwrite from MetadataOverwriteLM node
|
||||
|
||||
# Field names that the MetadataOverwriteLM node and its extractor share
|
||||
METADATA_OVERWRITE_FIELDS = (
|
||||
"prompt", "negative_prompt", "seed", "steps", "cfg_scale",
|
||||
"sampler", "scheduler", "model", "loras", "size",
|
||||
"clip_skip", "additional_data",
|
||||
)
|
||||
|
||||
# Complete list of categories to track
|
||||
METADATA_CATEGORIES = [MODELS, PROMPTS, SAMPLING, LORAS, EMBEDDINGS, SIZE, IMAGES]
|
||||
METADATA_CATEGORIES = [MODELS, PROMPTS, SAMPLING, LORAS, EMBEDDINGS, SIZE, IMAGES, OVERWRITE]
|
||||
|
||||
@@ -16,7 +16,7 @@ class MetadataHook:
|
||||
execution = None
|
||||
try:
|
||||
# Try direct import first
|
||||
import execution # type: ignore
|
||||
import execution # pyright: ignore[reportMissingImports]
|
||||
except ImportError:
|
||||
# Try to locate from system modules
|
||||
for module_name in sys.modules:
|
||||
@@ -83,7 +83,8 @@ class MetadataHook:
|
||||
|
||||
# Record inputs before execution
|
||||
if node_id is not None:
|
||||
registry.record_node_execution(node_id, class_type, input_data_all, None)
|
||||
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||
registry.record_node_execution(node_id, class_type, input_data_all, None, return_types=return_types)
|
||||
except Exception as e:
|
||||
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
||||
|
||||
@@ -114,7 +115,8 @@ class MetadataHook:
|
||||
|
||||
# Record outputs after execution
|
||||
if node_id is not None:
|
||||
registry.update_node_execution(node_id, class_type, results)
|
||||
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||
registry.update_node_execution(node_id, class_type, results, return_types=return_types)
|
||||
except Exception as e:
|
||||
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
||||
|
||||
@@ -136,6 +138,9 @@ class MetadataHook:
|
||||
if hasattr(prompt, 'original_prompt'):
|
||||
registry.set_current_prompt(prompt)
|
||||
|
||||
# Store extra_data for accessing full workflow node properties
|
||||
registry.set_extra_data(extra_data)
|
||||
|
||||
# Execute the original function
|
||||
return original_execute(*args, **kwargs)
|
||||
|
||||
@@ -163,7 +168,8 @@ class MetadataHook:
|
||||
class_type = obj.__class__.__name__
|
||||
node_id = unique_id
|
||||
if node_id is not None:
|
||||
registry.record_node_execution(node_id, class_type, input_data_all, None)
|
||||
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||
registry.record_node_execution(node_id, class_type, input_data_all, None, return_types=return_types)
|
||||
except Exception as e:
|
||||
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
||||
|
||||
@@ -180,7 +186,8 @@ class MetadataHook:
|
||||
class_type = obj.__class__.__name__
|
||||
node_id = unique_id
|
||||
if node_id is not None:
|
||||
registry.update_node_execution(node_id, class_type, results)
|
||||
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||
registry.update_node_execution(node_id, class_type, results, return_types=return_types)
|
||||
except Exception as e:
|
||||
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
||||
|
||||
@@ -202,6 +209,9 @@ class MetadataHook:
|
||||
if hasattr(prompt, 'original_prompt'):
|
||||
registry.set_current_prompt(prompt)
|
||||
|
||||
# Store extra_data for accessing full workflow node properties
|
||||
registry.set_extra_data(extra_data)
|
||||
|
||||
# Execute the original function
|
||||
return await original_execute(*args, **kwargs)
|
||||
|
||||
|
||||
@@ -1,15 +1,68 @@
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from .constants import IMAGES
|
||||
|
||||
# Check if running in standalone mode
|
||||
standalone_mode = os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1" or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0"
|
||||
|
||||
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IS_SAMPLER
|
||||
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IS_SAMPLER, OVERWRITE
|
||||
from .node_extractors import NODE_EXTRACTORS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Keys that identify metadata hint marks stored in node.properties.lm_marker_role
|
||||
_META_MARK_PREFIX = "meta_"
|
||||
_MARK_PRIMARY_MODEL = "primary_model"
|
||||
_MARK_PRIMARY_SAMPLER = "primary_sampler"
|
||||
_MARK_POSITIVE_PROMPT = "positive_prompt"
|
||||
_MARK_NEGATIVE_PROMPT = "negative_prompt"
|
||||
|
||||
class MetadataProcessor:
|
||||
"""Process and format collected metadata"""
|
||||
|
||||
@staticmethod
|
||||
def _get_user_marks(metadata):
|
||||
"""Scan workflow nodes (from extra_data.extra_pnginfo.workflow) for user-assigned
|
||||
metadata hint marks stored in node.properties.lm_marker_role.
|
||||
|
||||
Returns a dict mapping mark type keys to node IDs.
|
||||
Example: {'primary_model': '42', 'primary_sampler': '17'}
|
||||
"""
|
||||
marks: dict[str, str] = {}
|
||||
|
||||
# Primary source: extra_data.extra_pnginfo.workflow.nodes (has full properties)
|
||||
extra_data = metadata.get("extra_data")
|
||||
if extra_data and isinstance(extra_data, dict):
|
||||
extra_pnginfo = extra_data.get("extra_pnginfo", {})
|
||||
if isinstance(extra_pnginfo, dict):
|
||||
workflow = extra_pnginfo.get("workflow", {})
|
||||
nodes = workflow.get("nodes", [])
|
||||
for node in nodes:
|
||||
node_id = str(node.get("id", ""))
|
||||
role = node.get("properties", {}).get("lm_marker_role", "")
|
||||
if role.startswith(_META_MARK_PREFIX):
|
||||
mark_type = role[len(_META_MARK_PREFIX):]
|
||||
if mark_type in marks:
|
||||
logger.warning(
|
||||
"Duplicate meta hint '%s': node %s (previous: %s), "
|
||||
"last match wins",
|
||||
mark_type, node_id, marks[mark_type],
|
||||
)
|
||||
marks[mark_type] = node_id
|
||||
|
||||
# Fallback: try prompt.original_prompt (API-only submissions may not have workflow)
|
||||
if not marks:
|
||||
prompt = metadata.get("current_prompt")
|
||||
if prompt and getattr(prompt, "original_prompt", None):
|
||||
for node_id, node_data in prompt.original_prompt.items():
|
||||
role = node_data.get("properties", {}).get("lm_marker_role", "")
|
||||
if role.startswith(_META_MARK_PREFIX):
|
||||
mark_type = role[len(_META_MARK_PREFIX):]
|
||||
marks[mark_type] = node_id
|
||||
|
||||
return marks
|
||||
|
||||
@staticmethod
|
||||
def find_primary_sampler(metadata, downstream_id=None):
|
||||
"""
|
||||
@@ -162,6 +215,24 @@ class MetadataProcessor:
|
||||
primary_sampler = sampler_info
|
||||
primary_sampler_id = node_id
|
||||
|
||||
# Last resort: any registered sampler. Samplers without a denoise or
|
||||
# add_noise parameter (e.g. multi-stage samplers like KreaTwoStageSampler)
|
||||
# are not caught by the criteria above. Prefer execution order so the
|
||||
# first executed sampler wins, matching the downstream_id branch.
|
||||
if primary_sampler is None:
|
||||
sampler_ids = [
|
||||
node_id
|
||||
for node_id, sampler_info in metadata.get(SAMPLING, {}).items()
|
||||
if sampler_info.get(IS_SAMPLER, False)
|
||||
]
|
||||
if sampler_ids:
|
||||
if downstream_id and "execution_order" in metadata:
|
||||
for node_id in metadata["execution_order"]:
|
||||
if node_id in sampler_ids:
|
||||
return node_id, metadata[SAMPLING][node_id]
|
||||
primary_sampler_id = sampler_ids[0]
|
||||
primary_sampler = metadata[SAMPLING][sampler_ids[0]]
|
||||
|
||||
return primary_sampler_id, primary_sampler
|
||||
|
||||
@staticmethod
|
||||
@@ -471,17 +542,54 @@ class MetadataProcessor:
|
||||
"checkpoint": None,
|
||||
"loras": "",
|
||||
"size": None,
|
||||
"clip_skip": None
|
||||
"clip_skip": None,
|
||||
"additional_data": "",
|
||||
}
|
||||
|
||||
# Get the prompt object for node relationship tracing
|
||||
prompt = metadata.get("current_prompt")
|
||||
|
||||
# Find the primary KSampler node
|
||||
# ---- User marks: override heuristic inference with user-assigned hints ----
|
||||
user_marks = MetadataProcessor._get_user_marks(metadata)
|
||||
|
||||
# Find the primary KSampler node (user mark takes priority)
|
||||
primary_sampler_id = None
|
||||
primary_sampler = None
|
||||
if _MARK_PRIMARY_SAMPLER in user_marks:
|
||||
marked_id = user_marks[_MARK_PRIMARY_SAMPLER]
|
||||
sampler_data = metadata.get(SAMPLING, {}).get(marked_id)
|
||||
if sampler_data and sampler_data.get(IS_SAMPLER):
|
||||
primary_sampler_id = marked_id
|
||||
primary_sampler = sampler_data
|
||||
else:
|
||||
logger.warning(
|
||||
"User-marked primary sampler %s has no runtime metadata, "
|
||||
"falling back to heuristic",
|
||||
marked_id,
|
||||
)
|
||||
if primary_sampler is None:
|
||||
primary_sampler_id, primary_sampler = MetadataProcessor.find_primary_sampler(metadata, id)
|
||||
|
||||
# Directly get checkpoint from metadata instead of tracing
|
||||
# Pass primary_sampler_id to avoid redundant calculation
|
||||
# Resolve checkpoint / model (user mark takes priority)
|
||||
if _MARK_PRIMARY_MODEL in user_marks:
|
||||
marked_id = user_marks[_MARK_PRIMARY_MODEL]
|
||||
if marked_id in metadata.get(MODELS, {}):
|
||||
params["checkpoint"] = metadata[MODELS][marked_id].get("name")
|
||||
else:
|
||||
extra_data = metadata.get("extra_data")
|
||||
extra_pnginfo = extra_data.get("extra_pnginfo", {}) if extra_data and isinstance(extra_data, dict) else {}
|
||||
workflow = extra_pnginfo.get("workflow", {}) if isinstance(extra_pnginfo, dict) else {}
|
||||
node_type = "unknown"
|
||||
for n in workflow.get("nodes", []):
|
||||
if str(n.get("id", "")) == marked_id:
|
||||
node_type = n.get("type", "unknown")
|
||||
break
|
||||
logger.warning(
|
||||
"User-marked primary model %s (type=%s, registered=%s) has no runtime metadata, "
|
||||
"falling back to heuristic",
|
||||
marked_id, node_type, node_type in NODE_EXTRACTORS,
|
||||
)
|
||||
if params["checkpoint"] is None:
|
||||
checkpoint = MetadataProcessor.find_primary_checkpoint(metadata, id, primary_sampler_id)
|
||||
if checkpoint:
|
||||
params["checkpoint"] = checkpoint
|
||||
@@ -540,6 +648,21 @@ class MetadataProcessor:
|
||||
# For SamplerCustom, handle any additional parameters
|
||||
MetadataProcessor.handle_custom_advanced_sampler(metadata, prompt, primary_sampler_id, params)
|
||||
|
||||
# ---- User marks: override prompts with explicitly tagged nodes ----
|
||||
prompts_data = metadata.get(PROMPTS, {})
|
||||
if _MARK_POSITIVE_PROMPT in user_marks:
|
||||
pos_id = user_marks[_MARK_POSITIVE_PROMPT]
|
||||
if pos_id in prompts_data:
|
||||
prompt_text = prompts_data[pos_id].get("text") or prompts_data[pos_id].get("positive_text")
|
||||
if prompt_text:
|
||||
params["prompt"] = prompt_text
|
||||
if _MARK_NEGATIVE_PROMPT in user_marks:
|
||||
neg_id = user_marks[_MARK_NEGATIVE_PROMPT]
|
||||
if neg_id in prompts_data:
|
||||
prompt_text = prompts_data[neg_id].get("text") or prompts_data[neg_id].get("negative_text")
|
||||
if prompt_text:
|
||||
params["negative_prompt"] = prompt_text
|
||||
|
||||
# Size extraction is same for all sampler types
|
||||
# Check if the sampler itself has size information (from latent_image)
|
||||
if primary_sampler_id in metadata.get(SIZE, {}):
|
||||
@@ -569,6 +692,25 @@ class MetadataProcessor:
|
||||
if params["clip_skip"] is None:
|
||||
params["clip_skip"] = "1"
|
||||
|
||||
# ---- Apply manual metadata overwrites ----
|
||||
for overwrite_info in metadata.get(OVERWRITE, {}).values():
|
||||
overwrite_params = overwrite_info.get("parameters", {})
|
||||
for key, value in overwrite_params.items():
|
||||
if key == "clip_skip":
|
||||
# Accept any value from overwrite node (sentinel -25 already
|
||||
# filtered upstream). Needed because falsy check treats 0
|
||||
# as "not set" even though 0 is a valid wired input here.
|
||||
params[key] = value
|
||||
elif value: # truthy check — only overwrite when user provided a real value
|
||||
params[key] = value
|
||||
|
||||
# Bridge: the overwrite node exposes the field as "model" (more accurate),
|
||||
# but the internal pipeline key remains "checkpoint" for backward compatibility
|
||||
# with A1111 metadata format and downstream consumers.
|
||||
if params.get("model"):
|
||||
params["checkpoint"] = params["model"]
|
||||
del params["model"]
|
||||
|
||||
return params
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
import time
|
||||
from nodes import NODE_CLASS_MAPPINGS # type: ignore
|
||||
from typing import Any
|
||||
from nodes import NODE_CLASS_MAPPINGS # pyright: ignore[reportMissingImports, reportAttributeAccessIssue]
|
||||
from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor
|
||||
from .constants import METADATA_CATEGORIES, IMAGES
|
||||
from .constants import METADATA_CATEGORIES, IMAGES, OVERWRITE
|
||||
|
||||
|
||||
class MetadataRegistry:
|
||||
@@ -9,6 +10,15 @@ class MetadataRegistry:
|
||||
|
||||
_instance = None
|
||||
|
||||
current_prompt_id: Any = None
|
||||
current_prompt: Any = None
|
||||
metadata: dict[str, Any] = {}
|
||||
prompt_metadata: dict[str, Any] = {}
|
||||
executed_nodes: set[str] = set()
|
||||
node_cache: dict[str, Any] = {}
|
||||
max_prompt_history: int = 3
|
||||
metadata_categories: list[str] = METADATA_CATEGORIES
|
||||
|
||||
def __new__(cls):
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
@@ -61,6 +71,7 @@ class MetadataRegistry:
|
||||
{
|
||||
"execution_order": [],
|
||||
"current_prompt": None, # Will store the prompt object
|
||||
"extra_data": None, # Will store the API extra_data for workflow metadata
|
||||
"timestamp": time.time(),
|
||||
}
|
||||
)
|
||||
@@ -75,6 +86,11 @@ class MetadataRegistry:
|
||||
# Store the prompt in the metadata for later relationship tracing
|
||||
self.prompt_metadata[self.current_prompt_id]["current_prompt"] = prompt
|
||||
|
||||
def set_extra_data(self, extra_data):
|
||||
"""Store the API extra_data (contains extra_pnginfo.workflow with node properties)"""
|
||||
if self.current_prompt_id and self.current_prompt_id in self.prompt_metadata:
|
||||
self.prompt_metadata[self.current_prompt_id]["extra_data"] = extra_data
|
||||
|
||||
def get_metadata(self, prompt_id=None):
|
||||
"""Get collected metadata for a prompt"""
|
||||
key = prompt_id if prompt_id is not None else self.current_prompt_id
|
||||
@@ -122,20 +138,28 @@ class MetadataRegistry:
|
||||
cache_key = f"{node_id}:{class_type}"
|
||||
|
||||
# Check if this node type is relevant for metadata collection
|
||||
if class_type in NODE_EXTRACTORS:
|
||||
if class_type in NODE_EXTRACTORS or cache_key in self.node_cache:
|
||||
# Check if we have cached metadata for this node
|
||||
if cache_key in self.node_cache:
|
||||
cached_data = self.node_cache[cache_key]
|
||||
|
||||
# Detect bypass (mode=4) / mute (mode=2) — these nodes
|
||||
# were intentionally disabled and should not contribute
|
||||
# overwrite values from a previous execution's cache.
|
||||
node_mode = node_data.get("mode", 0)
|
||||
node_is_disabled = node_mode in (2, 4)
|
||||
|
||||
# Apply cached metadata to the current metadata
|
||||
for category in self.metadata_categories:
|
||||
if category == OVERWRITE and node_is_disabled:
|
||||
continue
|
||||
if category in cached_data and node_id in cached_data[category]:
|
||||
if node_id not in metadata[category]:
|
||||
metadata[category][node_id] = cached_data[category][
|
||||
node_id
|
||||
]
|
||||
|
||||
def record_node_execution(self, node_id, class_type, inputs, outputs):
|
||||
def record_node_execution(self, node_id, class_type, inputs, outputs, return_types=None):
|
||||
"""Record information about a node's execution"""
|
||||
if not self.current_prompt_id:
|
||||
return
|
||||
@@ -158,17 +182,18 @@ class MetadataRegistry:
|
||||
|
||||
# Extract node-specific metadata
|
||||
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
||||
extractor.extract(
|
||||
node_id,
|
||||
processed_inputs,
|
||||
outputs,
|
||||
if extractor is GenericNodeExtractor:
|
||||
extractor.extract(node_id, processed_inputs, outputs,
|
||||
self.prompt_metadata[self.current_prompt_id],
|
||||
)
|
||||
return_types=return_types)
|
||||
else:
|
||||
extractor.extract(node_id, processed_inputs, outputs,
|
||||
self.prompt_metadata[self.current_prompt_id])
|
||||
|
||||
# Cache this node's metadata
|
||||
self._cache_node_metadata(node_id, class_type)
|
||||
|
||||
def update_node_execution(self, node_id, class_type, outputs):
|
||||
def update_node_execution(self, node_id, class_type, outputs, return_types=None):
|
||||
"""Update node metadata with output information"""
|
||||
if not self.current_prompt_id:
|
||||
return
|
||||
@@ -179,8 +204,16 @@ class MetadataRegistry:
|
||||
# Use the same extractor to update with outputs
|
||||
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
||||
if hasattr(extractor, "update"):
|
||||
if extractor is GenericNodeExtractor:
|
||||
extractor.update(
|
||||
node_id, processed_outputs, self.prompt_metadata[self.current_prompt_id]
|
||||
node_id, processed_outputs,
|
||||
self.prompt_metadata[self.current_prompt_id],
|
||||
return_types=return_types,
|
||||
)
|
||||
else:
|
||||
extractor.update(
|
||||
node_id, processed_outputs,
|
||||
self.prompt_metadata[self.current_prompt_id],
|
||||
)
|
||||
|
||||
# Update the cached metadata for this node
|
||||
|
||||
@@ -2,7 +2,8 @@ import json
|
||||
import os
|
||||
import re
|
||||
|
||||
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER
|
||||
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER, OVERWRITE
|
||||
from .overwrite_utils import collect_overwrite_params
|
||||
|
||||
|
||||
def _store_checkpoint_metadata(metadata, node_id, model_name):
|
||||
@@ -31,10 +32,94 @@ class NodeMetadataExtractor:
|
||||
pass
|
||||
|
||||
class GenericNodeExtractor(NodeMetadataExtractor):
|
||||
"""Default extractor for nodes without specific handling"""
|
||||
"""Fallback extractor with type-signature-based detection.
|
||||
|
||||
When a node is not in the NODE_EXTRACTORS registry, the hook layer
|
||||
passes ``return_types`` from ``obj.RETURN_TYPES``:
|
||||
|
||||
* ``MODEL`` output: common input fields (ckpt_name, unet_name, etc.)
|
||||
are checked for a model file name and stored as checkpoint metadata.
|
||||
* ``CONDITIONING`` output: common text input fields are checked for
|
||||
prompt text, and conditioning inputs are tracked through transforms.
|
||||
"""
|
||||
|
||||
# Input field names that carry a model path in loader-style nodes.
|
||||
_MODEL_NAME_FIELDS = (
|
||||
"ckpt_name", "unet_name", "model_path", "model_name", "gguf_name",
|
||||
)
|
||||
|
||||
# Extensions used by checkpoint_scanner.py — only record values that look
|
||||
# like real model filenames to avoid capturing unrelated string fields.
|
||||
_MODEL_EXTENSIONS = {
|
||||
".ckpt", ".pt", ".pt2", ".bin", ".pth", ".safetensors", ".pkl", ".sft", ".gguf",
|
||||
}
|
||||
|
||||
# Input field names that may carry prompt text in encoder-style nodes.
|
||||
_TEXT_FIELDS = ("text", "clip_l", "t5xxl", "prompt", "positive", "negative")
|
||||
|
||||
@staticmethod
|
||||
def extract(node_id, inputs, outputs, metadata):
|
||||
pass
|
||||
def extract(node_id, inputs, outputs, metadata, return_types=None):
|
||||
if return_types is None:
|
||||
return
|
||||
|
||||
# — MODEL loader detection (checkpoint / UNET / GGUF) —
|
||||
if "MODEL" in return_types or any("MODEL" in str(t) for t in return_types):
|
||||
for field in GenericNodeExtractor._MODEL_NAME_FIELDS:
|
||||
val = inputs.get(field)
|
||||
if val and isinstance(val, str) and val.strip():
|
||||
name = val.strip()
|
||||
if not any(name.lower().endswith(ext) for ext in GenericNodeExtractor._MODEL_EXTENSIONS):
|
||||
continue
|
||||
_store_checkpoint_metadata(metadata, node_id, name)
|
||||
return
|
||||
|
||||
# — CONDITIONING encoder / transform detection —
|
||||
if "CONDITIONING" in return_types or any("CONDITIONING" in str(t) for t in return_types):
|
||||
text = None
|
||||
for field in GenericNodeExtractor._TEXT_FIELDS:
|
||||
val = inputs.get(field)
|
||||
if val and isinstance(val, str) and val.strip():
|
||||
text = val.strip()
|
||||
break
|
||||
|
||||
input_conditionings = _collect_conditioning_inputs(inputs)
|
||||
if text or input_conditionings:
|
||||
prompt_metadata = _ensure_prompt_metadata(metadata, node_id)
|
||||
if text:
|
||||
prompt_metadata["text"] = text
|
||||
if input_conditionings:
|
||||
prompt_metadata["orig_conditionings"] = input_conditionings
|
||||
|
||||
@staticmethod
|
||||
def update(node_id, outputs, metadata, return_types=None):
|
||||
if return_types is None:
|
||||
return
|
||||
if "CONDITIONING" not in return_types and not any(
|
||||
"CONDITIONING" in str(t) for t in return_types
|
||||
):
|
||||
return
|
||||
if node_id not in metadata.get(PROMPTS, {}):
|
||||
return
|
||||
output_tuple = _first_output_tuple(outputs)
|
||||
if not output_tuple or len(output_tuple) < 1:
|
||||
return
|
||||
|
||||
conditioning_index = _first_conditioning_index(return_types)
|
||||
if conditioning_index is None or len(output_tuple) <= conditioning_index:
|
||||
return
|
||||
|
||||
output_conditioning = output_tuple[conditioning_index]
|
||||
if output_conditioning is None:
|
||||
return
|
||||
|
||||
prompt_metadata = metadata[PROMPTS][node_id]
|
||||
prompt_metadata["conditioning"] = output_conditioning
|
||||
_record_conditioning_source(
|
||||
metadata,
|
||||
node_id,
|
||||
output_conditioning,
|
||||
prompt_metadata.get("orig_conditionings", []),
|
||||
)
|
||||
|
||||
class CheckpointLoaderExtractor(NodeMetadataExtractor):
|
||||
@staticmethod
|
||||
@@ -349,6 +434,34 @@ def _first_output_tuple(outputs):
|
||||
return None
|
||||
|
||||
|
||||
def _first_conditioning_index(return_types):
|
||||
"""Return the index of the first CONDITIONING output slot, or None."""
|
||||
if not return_types:
|
||||
return None
|
||||
for index, return_type in enumerate(return_types):
|
||||
if "CONDITIONING" in str(return_type):
|
||||
return index
|
||||
return None
|
||||
|
||||
|
||||
def _collect_conditioning_inputs(inputs):
|
||||
"""Collect conditioning object inputs (``conditioning*`` keys).
|
||||
|
||||
Primitive values (None, str, int, float, bool) are excluded so scalar
|
||||
fields like ``conditioning_strength`` are not mistaken for conditioning
|
||||
objects during provenance tracking.
|
||||
"""
|
||||
if not inputs:
|
||||
return []
|
||||
return [
|
||||
value
|
||||
for input_name, value in inputs.items()
|
||||
if input_name.startswith("conditioning")
|
||||
and value is not None
|
||||
and not isinstance(value, (str, int, float, bool))
|
||||
]
|
||||
|
||||
|
||||
def _record_conditioning_source(
|
||||
metadata, node_id, output_conditioning, input_conditionings
|
||||
):
|
||||
@@ -361,6 +474,14 @@ def _record_conditioning_source(
|
||||
if not sources:
|
||||
return
|
||||
|
||||
# Identity-preserving selectors return one of their inputs unchanged:
|
||||
# only that input contributed to the output, so record it alone instead
|
||||
# of treating every input as a combination source.
|
||||
for conditioning in sources:
|
||||
if id(conditioning) == id(output_conditioning):
|
||||
sources = [conditioning]
|
||||
break
|
||||
|
||||
prompt_metadata = _ensure_prompt_metadata(metadata, node_id)
|
||||
prompt_metadata.setdefault("conditioning_sources", []).append(
|
||||
{
|
||||
@@ -440,13 +561,7 @@ class ConditioningCombineExtractor(NodeMetadataExtractor):
|
||||
if not inputs:
|
||||
return
|
||||
|
||||
input_conditionings = []
|
||||
for input_name in inputs:
|
||||
if (
|
||||
input_name.startswith("conditioning")
|
||||
and inputs[input_name] is not None
|
||||
):
|
||||
input_conditionings.append(inputs[input_name])
|
||||
input_conditionings = _collect_conditioning_inputs(inputs)
|
||||
|
||||
if input_conditionings:
|
||||
prompt_metadata = _ensure_prompt_metadata(metadata, node_id)
|
||||
@@ -746,6 +861,65 @@ class TSCKSamplerAdvancedExtractor(KSamplerAdvancedExtractor, TSCSamplerBaseExtr
|
||||
|
||||
# Update method is inherited from TSCSamplerBaseExtractor
|
||||
|
||||
class KreaTwoStageSamplerExtractor(BaseSamplerExtractor):
|
||||
"""Extractor for Krea Two/Three Stage Samplers (Auryg/Krea-2-Two-Stage-Sampler).
|
||||
|
||||
The node samples in two (or three) stages with per-stage settings
|
||||
(stage1_steps/stage2_steps, stage1_cfg/stage2_cfg, ...). The canonical
|
||||
metadata fields consumed by ``extract_generation_params`` (steps, cfg,
|
||||
sampler_name, scheduler) are derived from the base stage (stage 1; the
|
||||
three-stage variant reuses stage 1 settings for stage 3), while the full
|
||||
per-stage breakdown is preserved in the raw parameters.
|
||||
"""
|
||||
|
||||
# All per-stage parameter keys present on both node variants.
|
||||
_STAGE_PARAM_KEYS = (
|
||||
"stage1_steps", "stage1_cfg", "stage1_sampler_name", "stage1_scheduler",
|
||||
"stage2_steps", "stage2_cfg", "stage2_sampler_name", "stage2_scheduler",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def extract(node_id, inputs, outputs, metadata):
|
||||
if not inputs:
|
||||
return
|
||||
|
||||
BaseSamplerExtractor.extract_sampling_params(
|
||||
node_id,
|
||||
inputs,
|
||||
metadata,
|
||||
("seed", "handoff_percent", "stage3_handoff_percent")
|
||||
+ KreaTwoStageSamplerExtractor._STAGE_PARAM_KEYS,
|
||||
)
|
||||
|
||||
# Derive the canonical fields expected by extract_generation_params.
|
||||
sampling_params = metadata[SAMPLING][node_id]["parameters"]
|
||||
if "stage1_steps" in sampling_params or "stage2_steps" in sampling_params:
|
||||
sampling_params["steps"] = (
|
||||
(sampling_params.get("stage1_steps") or 0)
|
||||
+ (sampling_params.get("stage2_steps") or 0)
|
||||
)
|
||||
if "stage1_cfg" in sampling_params:
|
||||
sampling_params["cfg"] = sampling_params["stage1_cfg"]
|
||||
if "stage1_sampler_name" in sampling_params:
|
||||
sampling_params["sampler_name"] = sampling_params["stage1_sampler_name"]
|
||||
if "stage1_scheduler" in sampling_params:
|
||||
sampling_params["scheduler"] = sampling_params["stage1_scheduler"]
|
||||
|
||||
BaseSamplerExtractor.extract_conditioning(node_id, inputs, metadata)
|
||||
|
||||
# Prefer the final generation resolution; latent dims are the fallback.
|
||||
BaseSamplerExtractor.extract_latent_dimensions(node_id, inputs, metadata)
|
||||
final_width = inputs.get("final_width")
|
||||
final_height = inputs.get("final_height")
|
||||
if final_width and final_height:
|
||||
if SIZE not in metadata:
|
||||
metadata[SIZE] = {}
|
||||
metadata[SIZE][node_id] = {
|
||||
"width": final_width,
|
||||
"height": final_height,
|
||||
"node_id": node_id,
|
||||
}
|
||||
|
||||
class LoraLoaderExtractor(NodeMetadataExtractor):
|
||||
@staticmethod
|
||||
def extract(node_id, inputs, outputs, metadata):
|
||||
@@ -786,6 +960,37 @@ class ImageSizeExtractor(NodeMetadataExtractor):
|
||||
"node_id": node_id
|
||||
}
|
||||
|
||||
class KreaDualResolutionSelectorExtractor(NodeMetadataExtractor):
|
||||
"""Extract base resolution from Krea Dual Resolution Selector outputs
|
||||
(Auryg/Krea-2-Two-Stage-Sampler).
|
||||
|
||||
The node computes base/final dimensions at runtime from aspect ratio and
|
||||
megapixel settings, so the values are only available in the update phase
|
||||
(outputs: base_width, base_height, final_width, final_height, seed).
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def extract(node_id, inputs, outputs, metadata):
|
||||
# Dimensions are computed at runtime; nothing to do here.
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def update(node_id, outputs, metadata):
|
||||
output_tuple = _first_output_tuple(outputs)
|
||||
if not output_tuple or len(output_tuple) < 2:
|
||||
return
|
||||
width, height = output_tuple[0], output_tuple[1]
|
||||
if not isinstance(width, int) or not isinstance(height, int):
|
||||
return
|
||||
|
||||
if SIZE not in metadata:
|
||||
metadata[SIZE] = {}
|
||||
metadata[SIZE][node_id] = {
|
||||
"width": width,
|
||||
"height": height,
|
||||
"node_id": node_id,
|
||||
}
|
||||
|
||||
class RgthreePowerLoraLoaderExtractor(NodeMetadataExtractor):
|
||||
"""Extract LoRA metadata from rgthree Power Lora Loader.
|
||||
|
||||
@@ -1154,6 +1359,28 @@ class CR_ApplyControlNetStackExtractor(NodeMetadataExtractor):
|
||||
metadata[PROMPTS][node_id]["positive_encoded"] = transformed_positive
|
||||
metadata[PROMPTS][node_id]["negative_encoded"] = transformed_negative
|
||||
|
||||
class MetadataOverwriteExtractor(NodeMetadataExtractor):
|
||||
"""Extract manually specified metadata from MetadataOverwriteLM node.
|
||||
|
||||
Stores truthy input values under the OVERWRITE category so that
|
||||
extract_generation_params can merge them over the inferred params.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def extract(node_id, inputs, outputs, metadata):
|
||||
if not inputs:
|
||||
return
|
||||
|
||||
overwrite_params = collect_overwrite_params(inputs)
|
||||
|
||||
if overwrite_params:
|
||||
metadata.setdefault(OVERWRITE, {})
|
||||
metadata[OVERWRITE][node_id] = {
|
||||
"parameters": overwrite_params,
|
||||
"node_id": node_id,
|
||||
}
|
||||
|
||||
|
||||
# Registry of node-specific extractors
|
||||
# Keys are node class names
|
||||
NODE_EXTRACTORS = {
|
||||
@@ -1165,6 +1392,8 @@ NODE_EXTRACTORS = {
|
||||
"ClownsharKSampler_Beta": SamplerExtractor,
|
||||
"TSC_KSampler": TSCKSamplerExtractor, # Efficient Nodes
|
||||
"TSC_KSamplerAdvanced": TSCKSamplerAdvancedExtractor, # Efficient Nodes
|
||||
"KreaTwoStageSampler": KreaTwoStageSamplerExtractor, # Auryg/Krea-2-Two-Stage-Sampler
|
||||
"KreaThreeStageSampler": KreaTwoStageSamplerExtractor, # Auryg/Krea-2-Two-Stage-Sampler
|
||||
"KSamplerBasicPipe": KSamplerBasicPipeExtractor, # comfyui-impact-pack
|
||||
"KSamplerAdvancedBasicPipe": KSamplerAdvancedBasicPipeExtractor, # comfyui-impact-pack
|
||||
"KSampler_inspire_pipe": KSamplerBasicPipeExtractor, # comfyui-inspire-pack
|
||||
@@ -1216,10 +1445,13 @@ NODE_EXTRACTORS = {
|
||||
"GetNode": GetNodeExtractor,
|
||||
# Latent
|
||||
"EmptyLatentImage": ImageSizeExtractor,
|
||||
"KreaDualResolutionSelector": KreaDualResolutionSelectorExtractor, # Auryg/Krea-2-Two-Stage-Sampler
|
||||
# Flux
|
||||
"FluxGuidance": FluxGuidanceExtractor, # Add FluxGuidance
|
||||
"CFGGuider": CFGGuiderExtractor, # Add CFGGuider
|
||||
# Image
|
||||
"VAEDecode": VAEDecodeExtractor, # Added VAEDecode extractor
|
||||
# Metadata overwrite
|
||||
"MetadataOverwriteLM": MetadataOverwriteExtractor,
|
||||
# Add other nodes as needed
|
||||
}
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
"""Shared helpers for Metadata Overwrite node metadata collection.
|
||||
|
||||
Used by both the MetadataOverwriteLM node (execution time) and the
|
||||
MetadataOverwriteExtractor (hook time) so the conversion/filtering logic
|
||||
cannot drift between the two paths.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Any, Dict
|
||||
|
||||
from ..utils.utils import model_patcher_to_name, sampler_object_to_name
|
||||
from .constants import CLIP_SKIP_SENTINEL, METADATA_OVERWRITE_FIELDS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def collect_overwrite_params(values: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Convert node input values into non-default overwrite parameters.
|
||||
|
||||
For most fields, a falsy value (empty string, 0) means "not set" and is
|
||||
skipped. clip_skip uses a dedicated sentinel (-25) so that a wired value
|
||||
of 0 is preserved. The ``model`` field accepts either a manual string or
|
||||
a wired MODEL (ModelPatcher) connection; in the latter case the source
|
||||
model name is extracted from the patcher's ``cached_patcher_init`` and
|
||||
stored as a ComfyUI-style relative path. The ``sampler`` field likewise
|
||||
accepts a manual string or a wired SAMPLER (KSAMPLER) connection, from
|
||||
which the sampler name is extracted via the sampler function's name.
|
||||
"""
|
||||
result: Dict[str, Any] = {}
|
||||
for key in METADATA_OVERWRITE_FIELDS:
|
||||
value = values.get(key)
|
||||
if key == "model" and not isinstance(value, str):
|
||||
value = model_patcher_to_name(value)
|
||||
if value is None:
|
||||
logger.warning(
|
||||
"Could not extract model name from wired MODEL input "
|
||||
"(no cached_patcher_init); model metadata overwrite skipped"
|
||||
)
|
||||
elif key == "sampler" and not isinstance(value, str):
|
||||
value = sampler_object_to_name(value)
|
||||
if value is None:
|
||||
logger.warning(
|
||||
"Could not extract sampler name from wired SAMPLER input "
|
||||
"(unrecognized sampler function); sampler metadata overwrite skipped"
|
||||
)
|
||||
if key == "clip_skip":
|
||||
if value != CLIP_SKIP_SENTINEL:
|
||||
result[key] = value
|
||||
elif value:
|
||||
result[key] = value
|
||||
return result
|
||||
@@ -43,7 +43,7 @@ SCANNER_GETTER_NAMES = tuple(SCANNER_TYPE_MAP.keys())
|
||||
|
||||
async def _find_model_entry(
|
||||
model_path: str,
|
||||
) -> tuple[object, object, str | None] | tuple[None, None, None]:
|
||||
) -> tuple[Any, object, str | None] | tuple[None, None, None]:
|
||||
"""Iterate all scanners and return the first (scanner, entry, getter_name)
|
||||
that owns *model_path*. Returns ``(None, None, None)`` when no scanner
|
||||
claims it.
|
||||
@@ -73,7 +73,7 @@ async def _find_model_entry(
|
||||
|
||||
async def _find_scanner_for_model(
|
||||
model_path: str,
|
||||
) -> tuple[object, object] | tuple[None, None]:
|
||||
) -> tuple[Any, object] | tuple[None, None]:
|
||||
"""Find the (scanner, cache_entry) responsible for *model_path*."""
|
||||
scanner, entry, _ = await _find_model_entry(model_path)
|
||||
return scanner, entry
|
||||
|
||||
@@ -46,6 +46,16 @@ async def api_json_error(
|
||||
if request.path.startswith("/api/lm/previews") and exc.status == 404:
|
||||
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(
|
||||
"API %s %s returned HTTP %d: %s",
|
||||
request.method,
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
import logging
|
||||
from typing import List, Tuple
|
||||
import comfy.sd # type: ignore
|
||||
import folder_paths # type: ignore
|
||||
import os
|
||||
from typing import Any, List, 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__)
|
||||
@@ -12,20 +13,42 @@ class CheckpointLoaderLM:
|
||||
|
||||
Loads checkpoints from both standard ComfyUI folders and LoRA Manager's
|
||||
extra folder paths, providing a unified interface for checkpoint loading.
|
||||
The ckpt_name combo supports ComfyUI's control_after_generate, letting
|
||||
users pick a random checkpoint on every run; the base_model input narrows
|
||||
the random pool through a front-end extension that filters the combo
|
||||
options.
|
||||
"""
|
||||
|
||||
NAME = "Checkpoint Loader (LoraManager)"
|
||||
CATEGORY = "Lora Manager/loaders"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
# Get list of checkpoint names from scanner (includes extra folder paths)
|
||||
checkpoint_names = s._get_checkpoint_names()
|
||||
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."},
|
||||
{
|
||||
"tooltip": (
|
||||
"The name of the checkpoint (model) to load. Use "
|
||||
"control_after_generate to pick a random model on "
|
||||
"every run."
|
||||
),
|
||||
"control_after_generate": "fixed",
|
||||
},
|
||||
),
|
||||
"base_model": (
|
||||
base_models,
|
||||
{
|
||||
"default": "Any",
|
||||
"tooltip": (
|
||||
"Restrict the random selection pool to this base "
|
||||
"model. 'Any' uses the full pool."
|
||||
),
|
||||
},
|
||||
),
|
||||
}
|
||||
}
|
||||
@@ -58,7 +81,10 @@ class CheckpointLoaderLM:
|
||||
for item in cache.raw_data:
|
||||
if item.get("sub_type") == "checkpoint":
|
||||
file_path = item.get("file_path", "")
|
||||
if file_path:
|
||||
# Only offer models that still exist on disk so ComfyUI
|
||||
# flags missing checkpoints at queue time via
|
||||
# "value not in list" (the scanner cache can be stale).
|
||||
if file_path and os.path.exists(file_path):
|
||||
# Format using relative path with OS-native separator
|
||||
formatted_name = _format_model_name_for_comfyui(
|
||||
file_path, model_roots
|
||||
@@ -89,15 +115,68 @@ class CheckpointLoaderLM:
|
||||
logger.error(f"Error getting checkpoint names: {e}")
|
||||
return []
|
||||
|
||||
def load_checkpoint(self, ckpt_name: str) -> Tuple:
|
||||
@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"]
|
||||
|
||||
@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())
|
||||
|
||||
def load_checkpoint(
|
||||
self, ckpt_name: str, base_model: str = "Any"
|
||||
) -> Tuple[Any, Any, Any]:
|
||||
"""Load a checkpoint by name, supporting extra folder paths
|
||||
|
||||
Args:
|
||||
ckpt_name: The name of the checkpoint to load (relative path with extension)
|
||||
base_model: Only used by the front-end to filter the random pool
|
||||
|
||||
Returns:
|
||||
Tuple of (MODEL, CLIP, VAE)
|
||||
"""
|
||||
del base_model
|
||||
# Get absolute path from cache using ComfyUI-style name
|
||||
ckpt_path, metadata = get_checkpoint_info_absolute(ckpt_name)
|
||||
|
||||
|
||||
@@ -0,0 +1,124 @@
|
||||
"""Create Hook LoRA (LoraManager) — multi-LoRA hook node compatible with ComfyUI's built-in hook pipeline.
|
||||
|
||||
Produces ``("HOOKS",)`` output that chains seamlessly with downstream hook consumers
|
||||
(ConditioningSetProperties, SetHookKeyframes, CombineHooks, SetClipHooks, etc.).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
|
||||
from ..utils.utils import get_lora_info_absolute
|
||||
from .utils import (
|
||||
FlexibleOptionalInputType,
|
||||
any_type,
|
||||
apply_lora_syntax_format,
|
||||
get_loras_list,
|
||||
validate_lora_entries,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class CreateHookLoraLM:
|
||||
NAME = "Create Hook LoRA (LoraManager)"
|
||||
CATEGORY = "Lora Manager/hooks"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"text": (
|
||||
"AUTOCOMPLETE_TEXT_LORAS",
|
||||
{
|
||||
"placeholder": "Search LoRAs to add...",
|
||||
"tooltip": (
|
||||
"Search and select LoRAs. Each LoRA gets its own "
|
||||
"model/clip strength. Hooks chain with prev_hooks."
|
||||
),
|
||||
},
|
||||
),
|
||||
"loras": ("LORAS", {}),
|
||||
},
|
||||
"optional": FlexibleOptionalInputType(any_type),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, loras=None):
|
||||
"""Queue-time validation: reject missing local LoRAs before execution."""
|
||||
return validate_lora_entries({"loras": loras}) or True
|
||||
|
||||
RETURN_TYPES = ("HOOKS", "STRING", "STRING")
|
||||
RETURN_NAMES = ("HOOKS", "trigger_words", "active_loras")
|
||||
FUNCTION = "create_hook"
|
||||
|
||||
def create_hook(self, text: str, loras, **kwargs):
|
||||
"""Create a HookGroup from the selected LoRAs, chained with prev_hooks.
|
||||
|
||||
Each active LoRA from the widget is loaded and wrapped in a WeightHook
|
||||
via :func:`comfy.hooks.create_hook_lora`. All hooks are combined into a
|
||||
single group and returned alongside trigger words and a human-readable
|
||||
summary of the active LoRAs.
|
||||
"""
|
||||
del text # used by the frontend widget only
|
||||
|
||||
# Lazy imports: comfy is not available in CI/test environment at module level
|
||||
import comfy.hooks # pyright: ignore[reportMissingImports] # noqa: C0415
|
||||
import comfy.utils # pyright: ignore[reportMissingImports] # noqa: C0415
|
||||
|
||||
prev_hooks: comfy.hooks.HookGroup | None = kwargs.get("prev_hooks")
|
||||
|
||||
hook_group = prev_hooks.clone() if prev_hooks is not None else comfy.hooks.HookGroup()
|
||||
|
||||
all_trigger_words: list[str] = []
|
||||
active_loras: list[tuple[str, float, float]] = []
|
||||
|
||||
for lora in get_loras_list({"loras": loras}):
|
||||
if not lora.get("active", False):
|
||||
continue
|
||||
|
||||
lora_name = apply_lora_syntax_format(lora["name"])
|
||||
model_strength = float(lora["strength"])
|
||||
clip_strength = float(lora.get("clipStrength", model_strength))
|
||||
|
||||
# Skip useless no-op entries (both strengths are zero)
|
||||
if model_strength == 0.0 and clip_strength == 0.0:
|
||||
continue
|
||||
|
||||
lora_path, trigger_words = get_lora_info_absolute(lora_name)
|
||||
if not lora_path or not os.path.isfile(lora_path):
|
||||
logger.warning("LoRA '%s' not found — skipping", lora_name)
|
||||
continue
|
||||
|
||||
try:
|
||||
lora_weights = comfy.utils.load_torch_file(lora_path, safe_load=True)
|
||||
|
||||
lora_hooks = comfy.hooks.create_hook_lora(
|
||||
lora=lora_weights,
|
||||
strength_model=model_strength,
|
||||
strength_clip=clip_strength,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Failed to load LoRA '%s' — skipping", lora_name)
|
||||
continue
|
||||
hook_group = hook_group.clone_and_combine(lora_hooks)
|
||||
|
||||
active_loras.append((lora_name, model_strength, clip_strength))
|
||||
all_trigger_words.extend(trigger_words)
|
||||
|
||||
# Format trigger words (group mode separator)
|
||||
trigger_words_text = ",, ".join(all_trigger_words) if all_trigger_words else ""
|
||||
|
||||
# Format active LoRAs summary
|
||||
formatted_loras = []
|
||||
for name, model_s, clip_s in active_loras:
|
||||
if abs(model_s - clip_s) > 0.001:
|
||||
formatted_loras.append(
|
||||
f"<lora:{name}:{model_s}:{clip_s}>"
|
||||
)
|
||||
else:
|
||||
formatted_loras.append(f"<lora:{name}:{model_s}>")
|
||||
active_loras_text = " ".join(formatted_loras)
|
||||
|
||||
return (hook_group, trigger_words_text, active_loras_text)
|
||||
@@ -0,0 +1,45 @@
|
||||
"""Lora Info display node — pure frontend node for showing selected LoRA info.
|
||||
|
||||
This node does NOT participate in workflow execution. Its single optional
|
||||
"lora_source" input exists solely as a wire-connection anchor so that the
|
||||
frontend can traverse the graph and push selection data to connected info nodes.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class LoraInfoLM:
|
||||
"""Display node that shows filename and notes for the selected LoRA."""
|
||||
|
||||
NAME = "Lora Info (LoraManager)"
|
||||
CATEGORY = "Lora Manager/utils"
|
||||
DESCRIPTION = (
|
||||
"Displays information (filename, notes) about the currently selected "
|
||||
"LoRA. Connect any output from a LoRA Loader or Stacker to the "
|
||||
"lora_source input, then select a LoRA in the source widget — the "
|
||||
"info updates automatically. Does not affect workflow execution."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
RETURN_NAMES = ()
|
||||
OUTPUT_NODE = False
|
||||
FUNCTION = "noop"
|
||||
|
||||
def noop(self, **kwargs):
|
||||
# This node is display-only — no workflow execution needed.
|
||||
return ()
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
LoraInfoLM.NAME: LoraInfoLM,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
LoraInfoLM.NAME: "Lora Info (LoraManager)",
|
||||
}
|
||||
+16
-24
@@ -1,9 +1,8 @@
|
||||
import importlib
|
||||
import logging
|
||||
import re
|
||||
|
||||
import comfy.sd # type: ignore
|
||||
import comfy.utils # type: ignore
|
||||
import comfy.sd # pyright: ignore[reportMissingImports]
|
||||
import comfy.utils # pyright: ignore[reportMissingImports]
|
||||
|
||||
from ..utils.utils import get_lora_info_absolute
|
||||
from .utils import (
|
||||
@@ -14,6 +13,8 @@ from .utils import (
|
||||
extract_lora_name,
|
||||
get_loras_list,
|
||||
nunchaku_load_lora,
|
||||
parse_lora_syntax,
|
||||
validate_lora_entries,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -48,9 +49,9 @@ def _collect_stack_entries(lora_stack):
|
||||
return entries
|
||||
|
||||
|
||||
def _collect_widget_entries(kwargs):
|
||||
def _collect_widget_entries(loras):
|
||||
entries = []
|
||||
for lora in get_loras_list(kwargs):
|
||||
for lora in get_loras_list({"loras": loras}):
|
||||
if not lora.get("active", False):
|
||||
continue
|
||||
lora_name = apply_lora_syntax_format(lora["name"])
|
||||
@@ -138,20 +139,26 @@ class LoraLoaderLM:
|
||||
"placeholder": "Search LoRAs to add...",
|
||||
"tooltip": "Format: <lora:lora_name:strength> separated by spaces or punctuation",
|
||||
}),
|
||||
"loras": ("LORAS", {}),
|
||||
},
|
||||
"optional": FlexibleOptionalInputType(any_type),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, loras=None):
|
||||
"""Queue-time validation: reject missing local LoRAs before execution."""
|
||||
return validate_lora_entries({"loras": loras}) or True
|
||||
|
||||
RETURN_TYPES = ("MODEL", "CLIP", "STRING", "STRING")
|
||||
RETURN_NAMES = ("MODEL", "CLIP", "trigger_words", "loaded_loras")
|
||||
FUNCTION = "load_loras"
|
||||
|
||||
def load_loras(self, model, text, **kwargs):
|
||||
"""Loads multiple LoRAs based on the kwargs input and lora_stack."""
|
||||
def load_loras(self, model, text, loras, **kwargs):
|
||||
"""Loads multiple LoRAs based on the widget input and lora_stack."""
|
||||
del text
|
||||
clip = kwargs.get("clip", None)
|
||||
lora_entries = _collect_stack_entries(kwargs.get("lora_stack", None))
|
||||
lora_entries.extend(_collect_widget_entries(kwargs))
|
||||
lora_entries.extend(_collect_widget_entries(loras))
|
||||
|
||||
nunchaku_model_kind = detect_nunchaku_model_kind(model)
|
||||
if nunchaku_model_kind == "flux":
|
||||
@@ -189,25 +196,10 @@ class LoraTextLoaderLM:
|
||||
RETURN_NAMES = ("MODEL", "CLIP", "trigger_words", "loaded_loras")
|
||||
FUNCTION = "load_loras_from_text"
|
||||
|
||||
def parse_lora_syntax(self, text):
|
||||
"""Parse LoRA syntax from text input."""
|
||||
pattern = r"<lora:([^:>]+):([^:>]+)(?::([^:>]+))?>"
|
||||
matches = re.findall(pattern, text, re.IGNORECASE)
|
||||
|
||||
loras = []
|
||||
for match in matches:
|
||||
model_strength = float(match[1])
|
||||
loras.append({
|
||||
"name": match[0],
|
||||
"model_strength": model_strength,
|
||||
"clip_strength": float(match[2]) if match[2] else model_strength,
|
||||
})
|
||||
return loras
|
||||
|
||||
def load_loras_from_text(self, model, lora_syntax, clip=None, lora_stack=None):
|
||||
"""Load LoRAs based on text syntax input."""
|
||||
lora_entries = _collect_stack_entries(lora_stack)
|
||||
for lora in self.parse_lora_syntax(lora_syntax):
|
||||
for lora in parse_lora_syntax(lora_syntax):
|
||||
lora_path, trigger_words = get_lora_info_absolute(lora["name"])
|
||||
lora_entries.append({
|
||||
"name": lora["name"],
|
||||
|
||||
@@ -9,6 +9,7 @@ and tracks the last used combination for reuse.
|
||||
import logging
|
||||
import os
|
||||
from ..utils.utils import get_lora_info
|
||||
from .utils import validate_lora_entries
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -31,6 +32,11 @@ class LoraRandomizerLM:
|
||||
},
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, loras=None):
|
||||
"""Queue-time validation: reject missing local LoRAs before execution."""
|
||||
return validate_lora_entries({"loras": loras}) or True
|
||||
|
||||
RETURN_TYPES = ("LORA_STACK",)
|
||||
RETURN_NAMES = ("LORA_STACK",)
|
||||
|
||||
|
||||
@@ -1,26 +1,102 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
_STACK_INPUT_PATTERN = re.compile(r"^lora_stack(?:_([ab])|(\d+))$")
|
||||
|
||||
|
||||
def _is_stack_input(name: str) -> bool:
|
||||
return bool(_STACK_INPUT_PATTERN.match(name))
|
||||
|
||||
|
||||
def _stack_slot_number(name: str) -> int:
|
||||
"""Numeric slot used to order stack inputs; legacy a/b map to 1/2."""
|
||||
match = _STACK_INPUT_PATTERN.match(name)
|
||||
if not match:
|
||||
return -1
|
||||
letter, digits = match.group(1), match.group(2)
|
||||
if digits is not None:
|
||||
return int(digits)
|
||||
return 1 if letter == "a" else 2
|
||||
|
||||
|
||||
class _LoraStackOptionalInputs:
|
||||
"""Lookup that preserves explicit optional inputs and dynamic lora_stack slots."""
|
||||
|
||||
def __init__(self, explicit_inputs: dict[str, tuple[str, dict[str, Any]]]) -> None:
|
||||
self._explicit_inputs = explicit_inputs
|
||||
|
||||
def __contains__(self, item: object) -> bool:
|
||||
if not isinstance(item, str):
|
||||
return False
|
||||
return item in self._explicit_inputs or _is_stack_input(item)
|
||||
|
||||
def __getitem__(self, key: str) -> tuple[str, dict[str, Any]]:
|
||||
if key in self._explicit_inputs:
|
||||
return self._explicit_inputs[key]
|
||||
if _is_stack_input(key):
|
||||
return (
|
||||
"LORA_STACK",
|
||||
{
|
||||
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||
},
|
||||
)
|
||||
raise KeyError(key)
|
||||
|
||||
|
||||
class LoraStackCombinerLM:
|
||||
NAME = "Lora Stack Combiner (LoraManager)"
|
||||
CATEGORY = "Lora Manager/stackers"
|
||||
DESCRIPTION = (
|
||||
"Combines multiple LoRA stacks into a single stack. "
|
||||
"Supports dynamic inputs: connect a stack to add more inputs."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"lora_stack_a": ("LORA_STACK",),
|
||||
"lora_stack_b": ("LORA_STACK",),
|
||||
optional_inputs: dict[str, tuple[str, dict[str, Any]]] = {
|
||||
"lora_stack1": (
|
||||
"LORA_STACK",
|
||||
{
|
||||
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||
},
|
||||
),
|
||||
"lora_stack2": (
|
||||
"LORA_STACK",
|
||||
{
|
||||
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||
},
|
||||
),
|
||||
}
|
||||
|
||||
stack = inspect.stack()
|
||||
if len(stack) > 2 and stack[2].function == "get_input_info":
|
||||
optional_inputs = _LoraStackOptionalInputs(optional_inputs) # pyright: ignore[reportAssignmentType]
|
||||
|
||||
return {
|
||||
"required": {},
|
||||
"optional": optional_inputs,
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LORA_STACK",)
|
||||
RETURN_NAMES = ("LORA_STACK",)
|
||||
FUNCTION = "combine_stacks"
|
||||
|
||||
def combine_stacks(self, lora_stack_a, lora_stack_b):
|
||||
combined_stack = []
|
||||
def combine_stacks(self, lora_stack1=None, lora_stack2=None, **kwargs):
|
||||
stacks = {
|
||||
"lora_stack1": lora_stack1,
|
||||
"lora_stack2": lora_stack2,
|
||||
}
|
||||
for key, value in kwargs.items():
|
||||
if _is_stack_input(key) and value is not None:
|
||||
stacks[key] = value
|
||||
|
||||
if lora_stack_a:
|
||||
combined_stack.extend(lora_stack_a)
|
||||
if lora_stack_b:
|
||||
combined_stack.extend(lora_stack_b)
|
||||
combined_stack = []
|
||||
for key in sorted(stacks, key=_stack_slot_number):
|
||||
stack = stacks[key]
|
||||
if stack:
|
||||
combined_stack.extend(stack)
|
||||
|
||||
return (combined_stack,)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import os
|
||||
from ..utils.utils import get_lora_info
|
||||
from .utils import FlexibleOptionalInputType, any_type, apply_lora_syntax_format, extract_lora_name, get_loras_list
|
||||
from .utils import FlexibleOptionalInputType, any_type, apply_lora_syntax_format, extract_lora_name, get_loras_list, validate_lora_entries
|
||||
|
||||
import logging
|
||||
|
||||
@@ -18,16 +18,22 @@ class LoraStackerLM:
|
||||
"placeholder": "Search LoRAs to add...",
|
||||
"tooltip": "Format: <lora:lora_name:strength> separated by spaces or punctuation",
|
||||
}),
|
||||
"loras": ("LORAS", {}),
|
||||
},
|
||||
"optional": FlexibleOptionalInputType(any_type),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, loras=None):
|
||||
"""Queue-time validation: reject missing local LoRAs before execution."""
|
||||
return validate_lora_entries({"loras": loras}) or True
|
||||
|
||||
RETURN_TYPES = ("LORA_STACK", "STRING", "STRING")
|
||||
RETURN_NAMES = ("LORA_STACK", "trigger_words", "active_loras")
|
||||
FUNCTION = "stack_loras"
|
||||
|
||||
def stack_loras(self, text, **kwargs):
|
||||
"""Stacks multiple LoRAs based on the kwargs input without loading them."""
|
||||
def stack_loras(self, text, loras, **kwargs):
|
||||
"""Stacks multiple LoRAs based on the widget input without loading them."""
|
||||
stack = []
|
||||
active_loras = []
|
||||
all_trigger_words = []
|
||||
@@ -42,8 +48,8 @@ class LoraStackerLM:
|
||||
_, trigger_words = get_lora_info(lora_name)
|
||||
all_trigger_words.extend(trigger_words)
|
||||
|
||||
# Process loras from kwargs with support for both old and new formats
|
||||
loras_list = get_loras_list(kwargs)
|
||||
# Process loras from the widget with support for both old and new formats
|
||||
loras_list = get_loras_list({"loras": loras})
|
||||
for lora in loras_list:
|
||||
if not lora.get('active', False):
|
||||
continue
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
"""Node to resolve `<lora:name:strength>` syntax to absolute file system paths.
|
||||
|
||||
Takes the loaded_loras / active_loras STRING output from LoraLoaderLM or
|
||||
LoraStackerLM and resolves each lora name to its absolute path on disk via
|
||||
the scanner cache. Unknown names are returned as-is.
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from ..utils.utils import get_lora_info_absolute
|
||||
from .utils import parse_lora_syntax
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class LoraSyntaxToPath:
|
||||
NAME = "LoRA Syntax → Path (LoraManager)"
|
||||
CATEGORY = "Lora Manager/utils"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"lora_syntax": (
|
||||
"STRING",
|
||||
{
|
||||
"forceInput": True,
|
||||
"multiline": True,
|
||||
"tooltip": (
|
||||
"<lora:name:strength> formatted text from "
|
||||
"loaded_loras / active_loras output"
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("paths",)
|
||||
FUNCTION = "resolve"
|
||||
|
||||
def resolve(self, lora_syntax: str) -> tuple[str]:
|
||||
"""Parse <lora:...> syntax and resolve each name to its absolute path."""
|
||||
if not lora_syntax or not lora_syntax.strip():
|
||||
logger.info("Received empty lora_syntax input")
|
||||
return ("",)
|
||||
|
||||
parsed = parse_lora_syntax(lora_syntax)
|
||||
if not parsed:
|
||||
logger.info("No valid <lora:...> entries found in input")
|
||||
return ("",)
|
||||
|
||||
paths: list[str] = []
|
||||
for entry in parsed:
|
||||
try:
|
||||
absolute_path, _ = get_lora_info_absolute(entry["name"])
|
||||
paths.append(absolute_path)
|
||||
except Exception:
|
||||
logger.warning("Failed to resolve lora '%s', skipping", entry["name"])
|
||||
continue
|
||||
|
||||
return ("\n".join(paths),)
|
||||
@@ -0,0 +1,179 @@
|
||||
"""Metadata Overwrite node — allows users to manually specify generation parameters
|
||||
that override the automatically collected/inferred metadata.
|
||||
|
||||
Most inputs have falsy defaults (empty string / 0) which are skipped.
|
||||
clip_skip uses a sentinel default (-25) so that a wired value of 0 is
|
||||
preserved — both ComfyUI and A1111 conventions have no meaningful 0 value,
|
||||
but users may wire 0 to express "no clip skip / default".
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..metadata_collector.constants import CLIP_SKIP_SENTINEL as _CLIP_SKIP_SENTINEL
|
||||
from ..metadata_collector.overwrite_utils import collect_overwrite_params
|
||||
|
||||
|
||||
class MetadataOverwriteLM:
|
||||
NAME = "Metadata Overwrite (LoraManager)"
|
||||
CATEGORY = "Lora Manager/utils"
|
||||
DESCRIPTION = (
|
||||
"Manually specify generation parameters to override automatically collected "
|
||||
"metadata. Only filled/connected inputs will take effect — empty defaults "
|
||||
"are ignored."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, Any]:
|
||||
return {
|
||||
"optional": {
|
||||
"prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Positive prompt. Only overwrites when non-empty.",
|
||||
},
|
||||
),
|
||||
"negative_prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Negative prompt. Only overwrites when non-empty.",
|
||||
},
|
||||
),
|
||||
"seed": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 0xFFFFFFFFFFFFFFFF,
|
||||
"control_after_generate": False,
|
||||
"tooltip": "Seed value. Only overwrites when > 0.",
|
||||
},
|
||||
),
|
||||
"steps": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 10000,
|
||||
"tooltip": "Number of steps. Only overwrites when > 0.",
|
||||
},
|
||||
),
|
||||
"cfg_scale": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0,
|
||||
"min": 0.0,
|
||||
"max": 100.0,
|
||||
"tooltip": "CFG scale. Only overwrites when > 0.",
|
||||
},
|
||||
),
|
||||
"sampler": (
|
||||
"STRING,SAMPLER",
|
||||
{
|
||||
"default": "",
|
||||
"widgetType": "STRING",
|
||||
"tooltip": (
|
||||
"Sampler name. Fill in the name manually or "
|
||||
"connect a SAMPLER output (e.g. KSamplerSelect) "
|
||||
"— the sampler name is then extracted "
|
||||
"automatically. Note: ddim is recorded as "
|
||||
"euler (ComfyUI internal representation). "
|
||||
"Only overwrites when non-empty."
|
||||
),
|
||||
},
|
||||
),
|
||||
"scheduler": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Scheduler name. Only overwrites when non-empty.",
|
||||
},
|
||||
),
|
||||
"model": (
|
||||
"STRING,MODEL",
|
||||
{
|
||||
"default": "",
|
||||
"widgetType": "STRING",
|
||||
"tooltip": (
|
||||
"The checkpoint or diffusion model (UNet) used "
|
||||
"for generation. Fill in the name manually or "
|
||||
"connect a MODEL output — the model name is then "
|
||||
"extracted automatically. Only overwrites when "
|
||||
"non-empty."
|
||||
),
|
||||
},
|
||||
),
|
||||
"loras": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": (
|
||||
"LoRA syntax, e.g. <lora:name:strength> "
|
||||
"or <lora:name:model_strength:clip_strength>, "
|
||||
"separated by spaces. Only overwrites when non-empty."
|
||||
),
|
||||
},
|
||||
),
|
||||
"size": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": (
|
||||
"Image size in WIDTHxHEIGHT format (e.g. 512x768). "
|
||||
"Only overwrites when non-empty."
|
||||
),
|
||||
},
|
||||
),
|
||||
"clip_skip": (
|
||||
"INT",
|
||||
{
|
||||
"default": _CLIP_SKIP_SENTINEL,
|
||||
"min": -25,
|
||||
"max": 24,
|
||||
"tooltip": (
|
||||
"Clip skip (ComfyUI: -24..-1, A1111: 1+). "
|
||||
"Default -25 means not set — any other value "
|
||||
"overwrites."
|
||||
),
|
||||
},
|
||||
),
|
||||
"additional_data": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": (
|
||||
"Additional data to embed in the image metadata. "
|
||||
"Inserted between Clip skip and Model hash in the "
|
||||
"A1111-compatible parameters string. "
|
||||
'Example: "Copyright": "Some license info"'
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("METADATA",)
|
||||
RETURN_NAMES = ("metadata",)
|
||||
FUNCTION = "collect_metadata"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def collect_metadata(self, **kwargs: Any) -> tuple[dict[str, Any]]:
|
||||
"""Collect non-default input values into a metadata dict.
|
||||
|
||||
For most fields, a falsy value (empty string, 0) means "not set"
|
||||
and is skipped. clip_skip uses a dedicated sentinel (-25) so that
|
||||
a wired value of 0 is preserved and reaches the metadata pipeline.
|
||||
|
||||
The ``model`` field accepts either a manual string or a wired MODEL
|
||||
(ModelPatcher) connection; in the latter case the underlying model
|
||||
name is extracted from the patcher's ``cached_patcher_init`` and
|
||||
stored as a ComfyUI-style relative path. The ``sampler`` field
|
||||
likewise accepts a manual string or a wired SAMPLER (KSAMPLER)
|
||||
connection, from which the sampler name is extracted automatically.
|
||||
"""
|
||||
return (collect_overwrite_params(kwargs),)
|
||||
+12
-13
@@ -15,15 +15,15 @@ import os
|
||||
import re
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Optional, Tuple, Union
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
import comfy.utils # type: ignore
|
||||
import folder_paths # type: ignore
|
||||
import comfy.utils # pyright: ignore[reportMissingImports]
|
||||
import folder_paths # pyright: ignore[reportMissingImports]
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from safetensors import safe_open
|
||||
|
||||
from nunchaku.lora.flux.nunchaku_converter import (
|
||||
from nunchaku.lora.flux.nunchaku_converter import ( # pyright: ignore[reportMissingTypeStubs]
|
||||
pack_lowrank_weight,
|
||||
unpack_lowrank_weight,
|
||||
)
|
||||
@@ -87,10 +87,6 @@ def _rename_layer_underscore_layer_name(old_name: str) -> str:
|
||||
return new_name
|
||||
|
||||
|
||||
def _is_indexable_module(module):
|
||||
return isinstance(module, (nn.ModuleList, nn.Sequential, list, tuple))
|
||||
|
||||
|
||||
def _get_module_by_name(model: nn.Module, name: str) -> Optional[nn.Module]:
|
||||
if not name:
|
||||
return model
|
||||
@@ -100,7 +96,7 @@ def _get_module_by_name(model: nn.Module, name: str) -> Optional[nn.Module]:
|
||||
continue
|
||||
if hasattr(module, part):
|
||||
module = getattr(module, part)
|
||||
elif part.isdigit() and _is_indexable_module(module):
|
||||
elif part.isdigit() and isinstance(module, (nn.ModuleList, nn.Sequential, list, tuple)):
|
||||
try:
|
||||
module = module[int(part)]
|
||||
except (IndexError, TypeError):
|
||||
@@ -267,7 +263,9 @@ def _handle_proj_out_split(lora_dict: Dict[str, Dict[str, torch.Tensor]], base_k
|
||||
return result, consumed
|
||||
|
||||
|
||||
def _apply_lora_to_module(module: nn.Module, a_tensor: torch.Tensor, b_tensor: torch.Tensor, module_name: str, model: nn.Module) -> None:
|
||||
def _apply_lora_to_module(module: Any, a_tensor: torch.Tensor, b_tensor: torch.Tensor, module_name: str, model: Any) -> None:
|
||||
# These modules are dynamic torch containers; monkey-patched attributes
|
||||
# below are set at runtime, so the module/model types are deliberately Any.
|
||||
if not hasattr(module, "in_features") or not hasattr(module, "out_features"):
|
||||
raise ValueError(f"{module_name}: unsupported module without in/out features")
|
||||
if a_tensor.shape[1] != module.in_features or b_tensor.shape[0] != module.out_features:
|
||||
@@ -336,7 +334,7 @@ def _apply_lora_to_module(module: nn.Module, a_tensor: torch.Tensor, b_tensor: t
|
||||
raise ValueError(f"{module_name}: unsupported module type {type(module)}")
|
||||
|
||||
|
||||
def reset_lora_v2(model: nn.Module) -> None:
|
||||
def reset_lora_v2(model: Any) -> None:
|
||||
slots = getattr(model, "_lora_slots", None)
|
||||
if not slots:
|
||||
return
|
||||
@@ -344,6 +342,7 @@ def reset_lora_v2(model: nn.Module) -> None:
|
||||
module = _get_module_by_name(model, name)
|
||||
if module is None:
|
||||
continue
|
||||
module = cast(Any, module)
|
||||
module_type = info.get("type", "nunchaku")
|
||||
if module_type == "nunchaku":
|
||||
base_rank = info["base_rank"]
|
||||
@@ -371,7 +370,7 @@ def reset_lora_v2(model: nn.Module) -> None:
|
||||
def compose_loras_v2(model: nn.Module, lora_configs: List[Tuple[Union[str, Path, Dict[str, torch.Tensor]], float]], apply_awq_mod: bool = True) -> bool:
|
||||
del apply_awq_mod # retained for interface compatibility
|
||||
reset_lora_v2(model)
|
||||
aggregated_weights: Dict[str, List[Dict[str, object]]] = defaultdict(list)
|
||||
aggregated_weights: Dict[str, List[Dict[str, Any]]] = defaultdict(list)
|
||||
saw_supported_format = False
|
||||
unresolved_targets = 0
|
||||
|
||||
@@ -471,7 +470,7 @@ def compose_loras_v2(model: nn.Module, lora_configs: List[Tuple[Union[str, Path,
|
||||
class ComfyQwenImageWrapperLM(nn.Module):
|
||||
def __init__(self, model: nn.Module, config=None, apply_awq_mod: bool = True):
|
||||
super().__init__()
|
||||
self.model = model
|
||||
self.model: Any = model
|
||||
self.config = {} if config is None else config
|
||||
self.dtype = next(model.parameters()).dtype
|
||||
self.loras: List[Tuple[Union[str, Path, Dict[str, torch.Tensor]], float]] = []
|
||||
|
||||
+2
-2
@@ -67,7 +67,7 @@ class PromptLM:
|
||||
|
||||
stack = inspect.stack()
|
||||
if len(stack) > 2 and stack[2].function == "get_input_info":
|
||||
optional_inputs = _PromptOptionalInputs(optional_inputs) # type: ignore[assignment]
|
||||
optional_inputs = _PromptOptionalInputs(optional_inputs) # pyright: ignore[reportAssignmentType]
|
||||
|
||||
return {
|
||||
"required": {
|
||||
@@ -126,7 +126,7 @@ class PromptLM:
|
||||
else:
|
||||
prompt = expanded_text
|
||||
|
||||
from nodes import CLIPTextEncode # type: ignore
|
||||
from nodes import CLIPTextEncode # pyright: ignore[reportMissingImports, reportAttributeAccessIssue]
|
||||
|
||||
conditioning = CLIPTextEncode().encode(clip, prompt)[0]
|
||||
return (conditioning, prompt)
|
||||
|
||||
+360
-120
@@ -5,7 +5,7 @@ import time
|
||||
import uuid
|
||||
from typing import Any, Dict, Optional
|
||||
import numpy as np
|
||||
import folder_paths # type: ignore
|
||||
import folder_paths # pyright: ignore[reportMissingImports]
|
||||
from ..services.service_registry import ServiceRegistry
|
||||
from ..metadata_collector.metadata_processor import MetadataProcessor
|
||||
from ..metadata_collector import get_metadata
|
||||
@@ -13,9 +13,159 @@ from ..utils.constants import CARD_PREVIEW_WIDTH
|
||||
from ..utils.exif_utils import ExifUtils
|
||||
from ..utils.utils import calculate_recipe_fingerprint, sanitize_folder_name
|
||||
from PIL import Image, PngImagePlugin
|
||||
import piexif
|
||||
import piexif # pyright: ignore[reportMissingTypeStubs]
|
||||
import logging
|
||||
|
||||
# Civitai-compatible sampler name mapping: ComfyUI internal → A1111 display name
|
||||
CIVITAI_SAMPLER_MAP = {
|
||||
"euler": "Euler",
|
||||
"euler_ancestral": "Euler a",
|
||||
"lms": "LMS",
|
||||
"heun": "Heun",
|
||||
"dpm_2": "DPM2",
|
||||
"dpm_2_ancestral": "DPM2 a",
|
||||
"dpmpp_2s_ancestral": "DPM++ 2S a",
|
||||
"dpmpp_2m": "DPM++ 2M",
|
||||
"dpmpp_sde": "DPM++ SDE",
|
||||
"dpmpp_sde_gpu": "DPM++ SDE",
|
||||
"dpmpp_2m_sde": "DPM++ 2M SDE",
|
||||
"dpmpp_2m_sde_gpu": "DPM++ 2M SDE",
|
||||
"dpmpp_3m_sde": "DPM++ 3M SDE",
|
||||
"dpm_fast": "DPM fast",
|
||||
"dpm_adaptive": "DPM adaptive",
|
||||
"ddim": "DDIM",
|
||||
"plms": "PLMS",
|
||||
"uni_pc_bh2": "UniPC",
|
||||
"uni_pc": "UniPC",
|
||||
"lcm": "LCM",
|
||||
}
|
||||
|
||||
# Base model display name → AIR URN slug
|
||||
# Sourced from civitai source: src/shared/constants/basemodel.constants.ts
|
||||
BASE_MODEL_AIR_SLUG = {
|
||||
# Stable Diffusion family
|
||||
"SD 1.4": "sd1",
|
||||
"SD 1.5": "sd1",
|
||||
"SD 1.5 LCM": "sd1",
|
||||
"SD 1.5 Hyper": "sd1",
|
||||
"SD 2.0": "sd2",
|
||||
"SD 2.0 768": "sd2",
|
||||
"SD 2.1": "sd2",
|
||||
"SD 2.1 768": "sd2",
|
||||
"SD 2.1 Unclip": "sd2",
|
||||
"SD 3.0": "sd3",
|
||||
"SD 3.5": "sd35",
|
||||
"SD 3.5 Large": "sd35",
|
||||
"SD 3.5 Large Turbo": "sd35",
|
||||
"SD 3.5 Medium": "sd35",
|
||||
"SDXL 0.9": "sdxl",
|
||||
"SDXL 1.0": "sdxl",
|
||||
"SDXL 1.0 LCM": "sdxl",
|
||||
"SDXL Lightning": "sdxl",
|
||||
"SDXL Hyper": "sdxl",
|
||||
"SDXL Turbo": "sdxl",
|
||||
"SDXL Distilled": "sdxldistilled",
|
||||
"Stable Cascade": "scascade",
|
||||
"Stable Video Diffusion": "svd",
|
||||
"SVD": "svd",
|
||||
"SVD XT": "svdxt",
|
||||
|
||||
# SDXL community fine-tunes
|
||||
"Pony": "pony",
|
||||
"Pony Diffusion": "pony",
|
||||
"Illustrious": "illustrious",
|
||||
"NoobAI": "noobai",
|
||||
"Animagine": "illustrious",
|
||||
|
||||
# Flux family
|
||||
"Flux.1": "flux1",
|
||||
"Flux.1 D": "flux1",
|
||||
"Flux.1 S": "flux1",
|
||||
"Flux.1 Krea": "fluxkrea",
|
||||
"Flux.1 Kontext": "flux1kontext",
|
||||
"Flux.2": "flux2",
|
||||
"Flux.2 D": "flux2",
|
||||
"Flux.2 Klein 9B": "flux2klein_9b",
|
||||
"Flux.2 Klein 9B Base": "flux2klein_9b_base",
|
||||
"Flux.2 Klein 4B": "flux2klein_4b",
|
||||
"Flux.2 Klein 4B Base": "flux2klein_4b_base",
|
||||
|
||||
# Other image models (sorted alphabetically)
|
||||
"AuraFlow": "auraflow",
|
||||
"Chroma": "chroma",
|
||||
"HiDream": "hidream",
|
||||
"HiDream-O1": "hidream-o1",
|
||||
"Hunyuan DiT": "hydit1",
|
||||
"Hunyuan Video": "hyv1",
|
||||
"Kolors": "kolors",
|
||||
"Lumina": "lumina",
|
||||
"Mochi": "mochi",
|
||||
"ODOR": "odor",
|
||||
"PixArt Alpha": "pixarta",
|
||||
"PixArt Sigma": "pixarte",
|
||||
"Playground v2": "playgroundv2",
|
||||
"Playground v2.5": "playgroundv2",
|
||||
"Pony Diffusion V7": "ponyv7",
|
||||
|
||||
# Video models
|
||||
"CogVideoX": "cogvideox",
|
||||
"LTX Video": "ltxv",
|
||||
"LTX Video 2": "ltxv2",
|
||||
"LTX Video 2.3": "ltxv23",
|
||||
"Wan Video": "wanvideo",
|
||||
"Wan Video 1.3B T2V": "wanvideo_13b_t2v",
|
||||
"Wan Video 14B T2V": "wanvideo_14b_t2v",
|
||||
"Wan Video 14B I2V 480p": "wanvideo_14b_i2v_480p",
|
||||
"Wan Video 14B I2V 720p": "wanvideo_14b_i2v_720p",
|
||||
|
||||
# Third-party / proprietary image models
|
||||
"Boogu": "boogu",
|
||||
"Ernie": "ernie",
|
||||
"Grok": "grok",
|
||||
"HappyHorse": "happyhorse",
|
||||
"Ideogram": "ideogram",
|
||||
"Ideogram 4.0": "ideogram",
|
||||
"Imagen": "imagen4",
|
||||
"Imagen 4": "imagen4",
|
||||
"Krea": "krea2",
|
||||
"Krea 2": "krea2",
|
||||
"Lens": "lens",
|
||||
"MAI": "mai",
|
||||
"Nano Banana": "nanobanana",
|
||||
"OpenAI": "openai",
|
||||
"Reve": "reve",
|
||||
"Reve 2": "reve",
|
||||
"Reve 2.1": "reve",
|
||||
"Seedream": "seedream",
|
||||
"Sora": "sora2",
|
||||
"Sora 2": "sora2",
|
||||
"Veo": "veo3",
|
||||
"Veo 2": "veo3",
|
||||
"Veo 3": "veo3",
|
||||
"ZImageTurbo": "zimageturbo",
|
||||
"ZImageBase": "zimagebase",
|
||||
"ZImage": "zimagebase",
|
||||
|
||||
# Third-party video models
|
||||
"Hailuo by MiniMax": "minimax",
|
||||
"Haiper": "haiper",
|
||||
"Kling": "kling",
|
||||
"Lightricks": "lightricks",
|
||||
"Seedance": "seedance",
|
||||
"Vidu": "vidu",
|
||||
|
||||
# Qwen family
|
||||
"Qwen": "qwen",
|
||||
"Qwen 2": "qwen2",
|
||||
|
||||
# Anima
|
||||
"Anima": "anima",
|
||||
|
||||
# Special
|
||||
"Upscaler": "upscaler",
|
||||
"Other": "other",
|
||||
}
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -70,11 +220,29 @@ class SaveImageLM:
|
||||
"tooltip": "Compression quality for JPEG and lossy WebP formats (1-100). Higher values mean better quality but larger files.",
|
||||
},
|
||||
),
|
||||
"webp_method": (
|
||||
"INT",
|
||||
{
|
||||
"default": 6,
|
||||
"min": 0,
|
||||
"max": 6,
|
||||
"tooltip": "WebP compression method (0-6). 0=fastest/largest, 6=slowest/smallest. Only applies when file_format is 'webp'.",
|
||||
},
|
||||
),
|
||||
"jpeg_subsampling": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 2,
|
||||
"tooltip": "JPEG chroma subsampling level. 0=4:4:4 (best quality), 1=4:2:2, 2=4:2:0 (smallest files). Only applies when file_format is 'jpeg'.",
|
||||
},
|
||||
),
|
||||
"embed_workflow": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Embeds the complete workflow data into the image metadata. Only works with PNG and WebP formats.",
|
||||
"tooltip": "When enabled, saved images store the complete workflow. Drag the image back into ComfyUI to restore the original node graph. PNG and WebP only.",
|
||||
},
|
||||
),
|
||||
"save_with_metadata": (
|
||||
@@ -84,6 +252,13 @@ class SaveImageLM:
|
||||
"tooltip": "When enabled, embeds generation parameters into the saved image metadata. Disable to skip writing generation metadata.",
|
||||
},
|
||||
),
|
||||
"add_loras_to_prompt": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "When enabled, appends the LoRA syntax line (e.g. <lora:name:strength>) after the positive prompt in the saved metadata.",
|
||||
},
|
||||
),
|
||||
"add_counter_to_filename": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
@@ -142,148 +317,197 @@ class SaveImageLM:
|
||||
|
||||
return None
|
||||
|
||||
def format_metadata(self, metadata_dict):
|
||||
"""Format metadata in the requested format similar to userComment example"""
|
||||
if not metadata_dict:
|
||||
return ""
|
||||
def _resolve_model_cache_entry(self, scanner_type: str, name: str):
|
||||
"""Resolve model hash, civitai metadata, and base_model from scanner cache.
|
||||
Returns (hash_str, civitai_dict, base_model_str). All values are empty defaults when not found."""
|
||||
scanner = ServiceRegistry.get_service_sync(scanner_type)
|
||||
if scanner is None or not name:
|
||||
return "", {}, ""
|
||||
|
||||
# Helper function to only add parameter if value is not None
|
||||
def add_param_if_not_none(param_list, label, value):
|
||||
if value is not None:
|
||||
param_list.append(f"{label}: {value}")
|
||||
entry = self._get_cached_model_by_name(scanner, name)
|
||||
if entry is None:
|
||||
basename = os.path.splitext(os.path.basename(name))[0]
|
||||
hash_val = scanner.get_hash_by_filename(basename)
|
||||
return (hash_val or "").lower(), {}, ""
|
||||
|
||||
hash_val = (entry.get("sha256") or "").lower()
|
||||
civitai = entry.get("civitai") or {}
|
||||
base_model = entry.get("base_model") or ""
|
||||
return hash_val, civitai, base_model
|
||||
|
||||
@staticmethod
|
||||
def _get_civitai_sampler_name(sampler_name: str, scheduler: str) -> str:
|
||||
if sampler_name in CIVITAI_SAMPLER_MAP:
|
||||
civitai_name = CIVITAI_SAMPLER_MAP[sampler_name]
|
||||
if scheduler == "karras":
|
||||
civitai_name += " Karras"
|
||||
elif scheduler == "exponential":
|
||||
civitai_name += " Exponential"
|
||||
return civitai_name
|
||||
else:
|
||||
if scheduler and scheduler != "normal":
|
||||
return f"{sampler_name}_{scheduler}"
|
||||
return sampler_name
|
||||
|
||||
@staticmethod
|
||||
def _build_air_string(base_model: str, model_type: str, model_id: int, version_id: int) -> str:
|
||||
slug = BASE_MODEL_AIR_SLUG.get(base_model, "other")
|
||||
type_lower = model_type.lower() if model_type else "other"
|
||||
return f"urn:air:{slug}:{type_lower}:civitai:{model_id}@{version_id}"
|
||||
|
||||
def format_metadata(self, metadata_dict: dict[str, Any], add_loras_to_prompt: bool = False) -> str:
|
||||
"""Format metadata as A1111-compatible parameters string with Hashes JSON and Civitai resources."""
|
||||
if not metadata_dict: return ""
|
||||
|
||||
# Extract the prompt and negative prompt
|
||||
prompt = metadata_dict.get("prompt", "")
|
||||
negative_prompt = metadata_dict.get("negative_prompt", "")
|
||||
|
||||
# Extract loras from the prompt if present
|
||||
steps = metadata_dict.get("steps")
|
||||
cfg = metadata_dict.get("guidance")
|
||||
if cfg is None:
|
||||
cfg = metadata_dict.get("cfg_scale")
|
||||
if cfg is None:
|
||||
cfg = metadata_dict.get("cfg")
|
||||
seed = metadata_dict.get("seed")
|
||||
size = metadata_dict.get("size")
|
||||
sampler = metadata_dict.get("sampler") or ""
|
||||
scheduler = metadata_dict.get("scheduler") or "normal"
|
||||
checkpoint = metadata_dict.get("checkpoint") or ""
|
||||
loras_text = metadata_dict.get("loras", "")
|
||||
lora_hashes = {}
|
||||
clip_skip = metadata_dict.get("clip_skip")
|
||||
|
||||
# If loras are found, add them on a new line after the prompt
|
||||
# Parse LoRA entries from <lora:name:strength> format
|
||||
lora_entries: list[tuple[str, float]] = []
|
||||
if loras_text:
|
||||
prompt_with_loras = f"{prompt}\n{loras_text}"
|
||||
for match in re.findall(r"<lora:([^:]+):([^>]+)>", loras_text):
|
||||
lora_name, strength_str = match
|
||||
try:
|
||||
strength = float(strength_str)
|
||||
except (ValueError, TypeError):
|
||||
strength = 1.0
|
||||
lora_entries.append((lora_name, strength))
|
||||
|
||||
# Extract lora names from the format <lora:name:strength>
|
||||
lora_matches = re.findall(r"<lora:([^:]+):([^>]+)>", loras_text)
|
||||
# Resolve checkpoint hash and Civitai data from local cache
|
||||
ckpt_hash, ckpt_civitai, ckpt_base_model = "", {}, ""
|
||||
ckpt_display_name = ""
|
||||
if checkpoint:
|
||||
ckpt_hash, ckpt_civitai, ckpt_base_model = self._resolve_model_cache_entry(
|
||||
"checkpoint_scanner", checkpoint
|
||||
)
|
||||
ckpt_display_name = os.path.splitext(os.path.basename(checkpoint))[0]
|
||||
|
||||
# Get hash for each lora
|
||||
for lora_name, strength in lora_matches:
|
||||
hash_value = self.get_lora_hash(lora_name)
|
||||
if hash_value:
|
||||
lora_hashes[lora_name] = hash_value
|
||||
else:
|
||||
prompt_with_loras = prompt
|
||||
# Resolve LoRA hash and Civitai data from local cache
|
||||
loras_data: list[dict[str, Any]] = []
|
||||
for lora_name, strength in lora_entries:
|
||||
lora_hash, lora_civitai, lora_base_model = self._resolve_model_cache_entry(
|
||||
"lora_scanner", lora_name
|
||||
)
|
||||
loras_data.append({
|
||||
"name": lora_name,
|
||||
"strength": strength,
|
||||
"hash": lora_hash,
|
||||
"civitai": lora_civitai,
|
||||
"base_model": lora_base_model,
|
||||
})
|
||||
|
||||
# Format the first part (prompt and loras)
|
||||
metadata_parts = [prompt_with_loras]
|
||||
# Build Hashes JSON (A1111 / Civitai standard format)
|
||||
hashes: dict[str, str] = {}
|
||||
if ckpt_hash:
|
||||
hashes["model"] = ckpt_hash[:10].upper()
|
||||
for lora in loras_data:
|
||||
if lora["hash"]:
|
||||
hashes[f"LORA:{lora['name']}"] = lora["hash"][:10].upper()
|
||||
|
||||
# Add negative prompt
|
||||
if negative_prompt:
|
||||
metadata_parts.append(f"Negative prompt: {negative_prompt}")
|
||||
# Build Civitai resources JSON array
|
||||
civitai_resources: list[dict[str, Any]] = []
|
||||
if ckpt_civitai.get("id", 0) > 0:
|
||||
ckpt_resource: dict[str, Any] = {}
|
||||
ckpt_type = (ckpt_civitai.get("model") or {}).get("type", "Checkpoint")
|
||||
model_id = ckpt_civitai.get("modelId", 0)
|
||||
version_id = ckpt_civitai.get("id", 0)
|
||||
if model_id and version_id:
|
||||
ckpt_resource["air"] = self._build_air_string(
|
||||
ckpt_base_model, ckpt_type, int(model_id), int(version_id)
|
||||
)
|
||||
elif version_id:
|
||||
ckpt_resource["modelVersionId"] = int(version_id)
|
||||
if ckpt_civitai.get("name"):
|
||||
ckpt_resource["versionName"] = ckpt_civitai["name"]
|
||||
if ckpt_resource:
|
||||
civitai_resources.append(ckpt_resource)
|
||||
|
||||
# Format the second part (generation parameters)
|
||||
params = []
|
||||
for lora in loras_data:
|
||||
lora_civitai = lora["civitai"]
|
||||
if not lora_civitai or lora_civitai.get("id", 0) <= 0:
|
||||
continue
|
||||
lora_resource: dict[str, Any] = {"weight": lora["strength"]}
|
||||
lora_type = (lora_civitai.get("model") or {}).get("type", "LORA")
|
||||
model_id = lora_civitai.get("modelId", 0)
|
||||
version_id = lora_civitai.get("id", 0)
|
||||
if model_id and version_id:
|
||||
lora_resource["air"] = self._build_air_string(
|
||||
lora["base_model"], lora_type, int(model_id), int(version_id)
|
||||
)
|
||||
elif version_id:
|
||||
lora_resource["modelVersionId"] = int(version_id)
|
||||
if lora_civitai.get("name"):
|
||||
lora_resource["versionName"] = lora_civitai["name"]
|
||||
civitai_resources.append(lora_resource)
|
||||
|
||||
# Add standard parameters in the correct order
|
||||
if "steps" in metadata_dict:
|
||||
add_param_if_not_none(params, "Steps", metadata_dict.get("steps"))
|
||||
sampler_name = CIVITAI_SAMPLER_MAP.get(sampler, sampler) if sampler else None
|
||||
|
||||
# Combine sampler and scheduler information
|
||||
sampler_name = None
|
||||
scheduler_name = None
|
||||
|
||||
if "sampler" in metadata_dict:
|
||||
sampler = metadata_dict.get("sampler")
|
||||
# Convert ComfyUI sampler names to user-friendly names
|
||||
sampler_mapping = {
|
||||
"euler": "Euler",
|
||||
"euler_ancestral": "Euler a",
|
||||
"dpm_2": "DPM2",
|
||||
"dpm_2_ancestral": "DPM2 a",
|
||||
"heun": "Heun",
|
||||
"dpm_fast": "DPM fast",
|
||||
"dpm_adaptive": "DPM adaptive",
|
||||
"lms": "LMS",
|
||||
"dpmpp_2s_ancestral": "DPM++ 2S a",
|
||||
"dpmpp_sde": "DPM++ SDE",
|
||||
"dpmpp_sde_gpu": "DPM++ SDE",
|
||||
"dpmpp_2m": "DPM++ 2M",
|
||||
"dpmpp_2m_sde": "DPM++ 2M SDE",
|
||||
"dpmpp_2m_sde_gpu": "DPM++ 2M SDE",
|
||||
"ddim": "DDIM",
|
||||
}
|
||||
sampler_name = sampler_mapping.get(sampler, sampler)
|
||||
|
||||
if "scheduler" in metadata_dict:
|
||||
scheduler = metadata_dict.get("scheduler")
|
||||
scheduler_mapping = {
|
||||
"normal": "Simple",
|
||||
"normal": "Normal",
|
||||
"karras": "Karras",
|
||||
"exponential": "Exponential",
|
||||
"sgm_uniform": "SGM Uniform",
|
||||
"sgm_quadratic": "SGM Quadratic",
|
||||
}
|
||||
scheduler_name = scheduler_mapping.get(scheduler, scheduler)
|
||||
scheduler_name = scheduler_mapping.get(scheduler, scheduler) if scheduler else None
|
||||
|
||||
# Add combined sampler and scheduler information
|
||||
# Build output lines
|
||||
prompt_line = prompt if prompt else ""
|
||||
if add_loras_to_prompt and loras_text:
|
||||
prompt_line = f"{prompt_line}\n{loras_text}" if prompt_line else loras_text
|
||||
lines = [prompt_line] if prompt_line else [""]
|
||||
if negative_prompt:
|
||||
lines.append(f"Negative prompt: {negative_prompt}")
|
||||
|
||||
params: list[str] = []
|
||||
if steps is not None:
|
||||
params.append(f"Steps: {steps}")
|
||||
if sampler_name:
|
||||
if scheduler_name:
|
||||
params.append(f"Sampler: {sampler_name} {scheduler_name}")
|
||||
else:
|
||||
params.append(f"Sampler: {sampler_name}")
|
||||
|
||||
# CFG scale (Use guidance if available, otherwise fall back to cfg_scale or cfg)
|
||||
if "guidance" in metadata_dict:
|
||||
add_param_if_not_none(params, "CFG scale", metadata_dict.get("guidance"))
|
||||
elif "cfg_scale" in metadata_dict:
|
||||
add_param_if_not_none(params, "CFG scale", metadata_dict.get("cfg_scale"))
|
||||
elif "cfg" in metadata_dict:
|
||||
add_param_if_not_none(params, "CFG scale", metadata_dict.get("cfg"))
|
||||
|
||||
# Seed
|
||||
if "seed" in metadata_dict:
|
||||
add_param_if_not_none(params, "Seed", metadata_dict.get("seed"))
|
||||
|
||||
# Size
|
||||
if "size" in metadata_dict:
|
||||
add_param_if_not_none(params, "Size", metadata_dict.get("size"))
|
||||
|
||||
# Model info
|
||||
if "checkpoint" in metadata_dict:
|
||||
# Ensure checkpoint is a string before processing
|
||||
checkpoint = metadata_dict.get("checkpoint")
|
||||
if checkpoint is not None:
|
||||
# Get model hash
|
||||
model_hash = self.get_checkpoint_hash(checkpoint)
|
||||
|
||||
# Extract basename without path
|
||||
checkpoint_name = os.path.basename(checkpoint)
|
||||
# Remove extension if present
|
||||
checkpoint_name = os.path.splitext(checkpoint_name)[0]
|
||||
|
||||
# Add model hash if available
|
||||
if model_hash:
|
||||
if cfg is not None:
|
||||
params.append(f"CFG scale: {cfg}")
|
||||
if seed is not None:
|
||||
params.append(f"Seed: {seed}")
|
||||
if size:
|
||||
params.append(f"Size: {size}")
|
||||
if clip_skip is not None:
|
||||
try:
|
||||
params.append(f"Clip skip: {abs(int(clip_skip))}")
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
additional_data = metadata_dict.get("additional_data", "")
|
||||
if additional_data:
|
||||
params.append(additional_data)
|
||||
if ckpt_hash:
|
||||
params.append(f"Model hash: {ckpt_hash[:10].upper()}")
|
||||
if ckpt_display_name:
|
||||
params.append(f"Model: {ckpt_display_name}")
|
||||
if hashes:
|
||||
params.append(f"Hashes: {json.dumps(hashes, separators=(',', ':'))}")
|
||||
params.append("Version: ComfyUI")
|
||||
if civitai_resources:
|
||||
params.append(
|
||||
f"Model hash: {model_hash[:10]}, Model: {checkpoint_name}"
|
||||
f"Civitai resources: {json.dumps(civitai_resources, separators=(',', ':'))}"
|
||||
)
|
||||
else:
|
||||
params.append(f"Model: {checkpoint_name}")
|
||||
|
||||
# Add LoRA hashes if available
|
||||
if lora_hashes:
|
||||
lora_hash_parts = []
|
||||
for lora_name, hash_value in lora_hashes.items():
|
||||
lora_hash_parts.append(f"{lora_name}: {hash_value[:10]}")
|
||||
|
||||
if lora_hash_parts:
|
||||
params.append(f'Lora hashes: "{", ".join(lora_hash_parts)}"')
|
||||
|
||||
# Combine all parameters with commas
|
||||
metadata_parts.append(", ".join(params))
|
||||
|
||||
# Join all parts with a new line
|
||||
return "\n".join(metadata_parts)
|
||||
lines.append(", ".join(params))
|
||||
return "\n".join(lines)
|
||||
|
||||
# credit to nkchocoai
|
||||
# Add format_filename method to handle pattern substitution
|
||||
@@ -554,6 +778,14 @@ class SaveImageLM:
|
||||
if checkpoint_entry:
|
||||
recipe_data["checkpoint"] = checkpoint_entry
|
||||
|
||||
# The recipe image is the WebP produced above from the output file;
|
||||
# reuse the same metadata extraction to record workflow presence.
|
||||
try:
|
||||
metadata = ExifUtils._load_structured_metadata(image_path)
|
||||
recipe_data["has_workflow"] = bool(metadata.get("workflow"))
|
||||
except Exception:
|
||||
recipe_data["has_workflow"] = False
|
||||
|
||||
json_path = os.path.normpath(
|
||||
os.path.join(recipes_dir, f"{recipe_id}.recipe.json")
|
||||
)
|
||||
@@ -573,10 +805,13 @@ class SaveImageLM:
|
||||
extra_pnginfo=None,
|
||||
lossless_webp=True,
|
||||
quality=100,
|
||||
webp_method=6,
|
||||
jpeg_subsampling=0,
|
||||
embed_workflow=False,
|
||||
save_with_metadata=True,
|
||||
add_counter_to_filename=True,
|
||||
save_as_recipe=False,
|
||||
add_loras_to_prompt=False,
|
||||
):
|
||||
"""Save images with metadata"""
|
||||
results = []
|
||||
@@ -585,7 +820,7 @@ class SaveImageLM:
|
||||
raw_metadata = get_metadata()
|
||||
metadata_dict = MetadataProcessor.to_dict(raw_metadata, id)
|
||||
|
||||
metadata = self.format_metadata(metadata_dict)
|
||||
metadata = self.format_metadata(metadata_dict, add_loras_to_prompt)
|
||||
|
||||
# Process filename_prefix with pattern substitution
|
||||
filename_prefix = self.format_filename(filename_prefix, metadata_dict)
|
||||
@@ -627,15 +862,14 @@ class SaveImageLM:
|
||||
elif file_format == "jpeg":
|
||||
file = base_filename + ".jpg"
|
||||
file_extension = ".jpg"
|
||||
save_kwargs = {"quality": quality, "optimize": True}
|
||||
save_kwargs = {"quality": quality, "optimize": True, "subsampling": jpeg_subsampling}
|
||||
elif file_format == "webp":
|
||||
file = base_filename + ".webp"
|
||||
file_extension = ".webp"
|
||||
# Add optimization param to control performance
|
||||
save_kwargs = {
|
||||
"quality": quality,
|
||||
"lossless": lossless_webp,
|
||||
"method": 0,
|
||||
"method": webp_method,
|
||||
}
|
||||
else:
|
||||
raise ValueError(f"Unsupported file format: {file_format}")
|
||||
@@ -722,10 +956,13 @@ class SaveImageLM:
|
||||
extra_pnginfo=None,
|
||||
lossless_webp=True,
|
||||
quality=100,
|
||||
webp_method=6,
|
||||
jpeg_subsampling=0,
|
||||
embed_workflow=False,
|
||||
save_with_metadata=True,
|
||||
add_counter_to_filename=True,
|
||||
save_as_recipe=False,
|
||||
add_loras_to_prompt=False,
|
||||
):
|
||||
"""Process and save image with metadata"""
|
||||
# Make sure the output directory exists
|
||||
@@ -751,10 +988,13 @@ class SaveImageLM:
|
||||
extra_pnginfo,
|
||||
lossless_webp,
|
||||
quality,
|
||||
webp_method,
|
||||
jpeg_subsampling,
|
||||
embed_workflow,
|
||||
save_with_metadata,
|
||||
add_counter_to_filename,
|
||||
save_as_recipe,
|
||||
add_loras_to_prompt,
|
||||
)
|
||||
|
||||
return {
|
||||
|
||||
+107
-8
@@ -1,37 +1,74 @@
|
||||
import logging
|
||||
import os
|
||||
from typing import List, Tuple
|
||||
import comfy.sd # type: ignore
|
||||
from typing import Any, List, 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 UNETLoaderLM.load_unet so ModelPatcher
|
||||
deepclone/dynamic machinery can rebuild GGUF models with the correct
|
||||
GGMLOps. ``disable_dynamic`` is accepted for signature compatibility
|
||||
with core ComfyUI loaders.
|
||||
"""
|
||||
loader = UNETLoaderLM()
|
||||
model, = loader._load_gguf_unet(unet_path, unet_path, weight_dtype)
|
||||
return model
|
||||
|
||||
|
||||
class UNETLoaderLM:
|
||||
"""UNET Loader with support for extra folder paths
|
||||
|
||||
Loads diffusion models/UNets from both standard ComfyUI folders and LoRA Manager's
|
||||
extra folder paths, providing a unified interface for UNET loading.
|
||||
Supports both regular diffusion models and GGUF format models.
|
||||
The unet_name combo supports ComfyUI's control_after_generate, letting
|
||||
users pick a random diffusion model on every run; the base_model input
|
||||
narrows the random pool through a front-end extension that filters the
|
||||
combo options.
|
||||
"""
|
||||
|
||||
NAME = "Unet Loader (LoraManager)"
|
||||
CATEGORY = "Lora Manager/loaders"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
# Get list of unet names from scanner (includes extra folder paths)
|
||||
unet_names = s._get_unet_names()
|
||||
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."},
|
||||
{
|
||||
"tooltip": (
|
||||
"The name of the diffusion model to load. Use "
|
||||
"control_after_generate to pick a random model on "
|
||||
"every run."
|
||||
),
|
||||
"control_after_generate": "fixed",
|
||||
},
|
||||
),
|
||||
"weight_dtype": (
|
||||
["default", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2"],
|
||||
{"tooltip": "The dtype to use for the model weights."},
|
||||
),
|
||||
"base_model": (
|
||||
base_models,
|
||||
{
|
||||
"default": "Any",
|
||||
"tooltip": (
|
||||
"Restrict the random selection pool to this base "
|
||||
"model. 'Any' uses the full pool."
|
||||
),
|
||||
},
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -59,7 +96,10 @@ class UNETLoaderLM:
|
||||
for item in cache.raw_data:
|
||||
if item.get("sub_type") == "diffusion_model":
|
||||
file_path = item.get("file_path", "")
|
||||
if file_path:
|
||||
# Only offer models that still exist on disk so ComfyUI
|
||||
# flags missing diffusion models at queue time via
|
||||
# "value not in list" (the scanner cache can be stale).
|
||||
if file_path and os.path.exists(file_path):
|
||||
# Format using relative path with OS-native separator
|
||||
formatted_name = _format_model_name_for_comfyui(
|
||||
file_path, model_roots
|
||||
@@ -90,16 +130,69 @@ class UNETLoaderLM:
|
||||
logger.error(f"Error getting unet names: {e}")
|
||||
return []
|
||||
|
||||
def load_unet(self, unet_name: str, weight_dtype: str) -> Tuple:
|
||||
@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"]
|
||||
|
||||
@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())
|
||||
|
||||
def load_unet(
|
||||
self, unet_name: str, weight_dtype: str, 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
|
||||
base_model: Only used by the front-end to filter the random pool
|
||||
|
||||
Returns:
|
||||
Tuple of (MODEL,)
|
||||
"""
|
||||
del base_model
|
||||
import torch
|
||||
|
||||
# Get absolute path from cache using ComfyUI-style name
|
||||
@@ -133,7 +226,7 @@ class UNETLoaderLM:
|
||||
|
||||
def _load_gguf_unet(
|
||||
self, unet_path: str, unet_name: str, weight_dtype: str
|
||||
) -> Tuple:
|
||||
) -> Tuple[Any, ...]:
|
||||
"""Load a GGUF format diffusion model
|
||||
|
||||
Args:
|
||||
@@ -196,6 +289,12 @@ class UNETLoaderLM:
|
||||
# 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,)
|
||||
|
||||
except Exception as e:
|
||||
|
||||
+178
-2
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
class AnyType(str):
|
||||
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
|
||||
|
||||
@@ -6,7 +9,7 @@ class AnyType(str):
|
||||
|
||||
|
||||
# Credit to Regis Gaughan, III (rgthree)
|
||||
class FlexibleOptionalInputType(dict):
|
||||
class FlexibleOptionalInputType(dict[str, Any]):
|
||||
"""A special class to make flexible nodes that pass data to our python handlers.
|
||||
|
||||
Enables both flexible/dynamic input types (like for Any Switch) or a dynamic number of inputs
|
||||
@@ -23,6 +26,7 @@ class FlexibleOptionalInputType(dict):
|
||||
"""
|
||||
|
||||
def __init__(self, type):
|
||||
super().__init__()
|
||||
self.type = type
|
||||
|
||||
def __getitem__(self, key):
|
||||
@@ -36,10 +40,12 @@ any_type = AnyType("*")
|
||||
|
||||
# Common methods extracted from lora_loader.py and lora_stacker.py
|
||||
import os
|
||||
import re
|
||||
import logging
|
||||
import copy
|
||||
import sys
|
||||
import folder_paths # type: ignore
|
||||
import asyncio
|
||||
import folder_paths # pyright: ignore[reportMissingImports]
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -69,6 +75,25 @@ def extract_lora_name(lora_path):
|
||||
return apply_lora_syntax_format(name_no_ext)
|
||||
|
||||
|
||||
def parse_lora_syntax(text: str) -> list[dict[str, Any]]:
|
||||
"""Parse <lora:name:strength> syntax from text input into a list of dicts.
|
||||
|
||||
Each entry contains: name, model_strength, clip_strength.
|
||||
Supports both ``<lora:name:strength>`` and ``<lora:name:model_strength:clip_strength>``.
|
||||
"""
|
||||
pattern = r"<lora:([^:>]+):([^:>]+)(?::([^:>]+))?>"
|
||||
matches = re.findall(pattern, text, re.IGNORECASE)
|
||||
loras = []
|
||||
for match in matches:
|
||||
model_strength = float(match[1])
|
||||
loras.append({
|
||||
"name": match[0],
|
||||
"model_strength": model_strength,
|
||||
"clip_strength": float(match[2]) if match[2] else model_strength,
|
||||
})
|
||||
return loras
|
||||
|
||||
|
||||
def get_loras_list(kwargs):
|
||||
"""Helper to extract loras list from either old or new kwargs format"""
|
||||
if "loras" not in kwargs:
|
||||
@@ -87,6 +112,157 @@ def get_loras_list(kwargs):
|
||||
return []
|
||||
|
||||
|
||||
_LORA_EXTENSIONS = (".safetensors", ".ckpt", ".pt", ".bin")
|
||||
|
||||
|
||||
def _strip_lora_extension(name: str) -> str:
|
||||
"""Strip a known LoRA model extension from a name (case-insensitive)."""
|
||||
lowered = name.lower()
|
||||
for ext in _LORA_EXTENSIONS:
|
||||
if lowered.endswith(ext):
|
||||
return name[: -len(ext)]
|
||||
return name
|
||||
|
||||
|
||||
def _find_missing_loras(names: list[str]) -> list[str]:
|
||||
"""Return the names that cannot be resolved to an existing local LoRA file.
|
||||
|
||||
Mirrors the matching semantics of ``get_lora_info_absolute``
|
||||
(py/utils/utils.py): after stripping the extension, a name matches a cached
|
||||
LoRA when it equals the cached file name or the ``folder/file`` path. As a
|
||||
fallback, a name containing a folder that only matches by basename resolves
|
||||
to the first basename match (same behavior as the runtime resolver). Raw
|
||||
absolute paths that exist on disk are always considered available.
|
||||
|
||||
The scanner cache is fetched once for all names; the cache may be stale, so
|
||||
resolved paths are additionally verified with ``os.path.isfile``.
|
||||
"""
|
||||
if not names:
|
||||
return []
|
||||
|
||||
async def _check() -> list[str]:
|
||||
from ..services.service_registry import ServiceRegistry
|
||||
|
||||
scanner = await ServiceRegistry.get_lora_scanner()
|
||||
# The scanner cache may not be hydrated yet (startup, library path
|
||||
# change). An empty cache is not authoritative — treat it as "cannot
|
||||
# verify" and skip validation instead of flagging every active LoRA
|
||||
# as missing.
|
||||
if getattr(scanner, "_cache", None) is None or getattr(
|
||||
scanner, "_is_initializing", False
|
||||
):
|
||||
return []
|
||||
cache = await scanner.get_cached_data()
|
||||
|
||||
lookup = {}
|
||||
basename_candidates = {}
|
||||
for item in cache.raw_data:
|
||||
file_path = item.get("file_path")
|
||||
if not file_path:
|
||||
continue
|
||||
file_name = item.get("file_name", "")
|
||||
folder = item.get("folder", "")
|
||||
file_name_no_ext = _strip_lora_extension(file_name)
|
||||
path_name_no_ext = (
|
||||
f"{folder}/{file_name_no_ext}".replace("\\", "/")
|
||||
if folder
|
||||
else file_name_no_ext
|
||||
)
|
||||
lookup.setdefault(file_name_no_ext, file_path)
|
||||
lookup.setdefault(path_name_no_ext, file_path)
|
||||
basename_candidates.setdefault(file_name_no_ext, []).append(
|
||||
(folder, file_path)
|
||||
)
|
||||
|
||||
missing = []
|
||||
for name in names:
|
||||
if not name:
|
||||
continue
|
||||
normalized = name.replace("\\", "/")
|
||||
# Raw absolute paths (outside the library) are usable as-is.
|
||||
if os.path.isfile(normalized):
|
||||
continue
|
||||
no_ext = _strip_lora_extension(normalized)
|
||||
file_path = lookup.get(no_ext)
|
||||
if file_path is None and "/" in no_ext:
|
||||
# A name with a folder that matches only by basename resolves
|
||||
# at runtime like get_lora_info_absolute's fallback does:
|
||||
# prefer a candidate whose folder prefixes the name, else the
|
||||
# first basename match.
|
||||
folder, basename = no_ext.rsplit("/", 1)
|
||||
candidates = basename_candidates.get(basename, [])
|
||||
file_path = next(
|
||||
(
|
||||
fp
|
||||
for fld, fp in candidates
|
||||
if fld and no_ext.startswith(fld + "/")
|
||||
),
|
||||
None,
|
||||
)
|
||||
if file_path is None and candidates:
|
||||
file_path = candidates[0][1]
|
||||
if file_path is None or not os.path.isfile(file_path):
|
||||
missing.append(name)
|
||||
return missing
|
||||
|
||||
try:
|
||||
# Check if we're already in an event loop
|
||||
loop = asyncio.get_running_loop()
|
||||
# If we're in a running loop, run the async check in a separate thread
|
||||
import concurrent.futures
|
||||
|
||||
def run_in_thread():
|
||||
new_loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(new_loop)
|
||||
try:
|
||||
return new_loop.run_until_complete(_check())
|
||||
finally:
|
||||
new_loop.close()
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor() as executor:
|
||||
future = executor.submit(run_in_thread)
|
||||
return future.result()
|
||||
except RuntimeError:
|
||||
# No event loop is running, we can use asyncio.run()
|
||||
return asyncio.run(_check())
|
||||
|
||||
|
||||
def validate_lora_entries(kwargs):
|
||||
"""Validate active LoRA widget entries against the local library.
|
||||
|
||||
Used by node ``VALIDATE_INPUTS`` implementations so ComfyUI rejects the
|
||||
prompt at queue time (``custom_validation_failed``) when an active entry
|
||||
references a LoRA that is not available locally — mirroring how built-in
|
||||
loader nodes flag missing models before execution starts.
|
||||
|
||||
Returns:
|
||||
None when every active entry resolves to an existing local file,
|
||||
otherwise a descriptive error string listing the missing LoRAs.
|
||||
Verification failures (e.g. scanner not ready) are treated as valid
|
||||
so queueing is never blocked by validation machinery itself.
|
||||
"""
|
||||
# Missing/empty loras input is always valid; skip get_loras_list so it
|
||||
# does not log a warning for the None case on every queue.
|
||||
if not kwargs.get("loras"):
|
||||
return None
|
||||
loras = get_loras_list(kwargs)
|
||||
active_names = []
|
||||
for lora in loras:
|
||||
if not isinstance(lora, dict):
|
||||
continue
|
||||
if not lora.get("active", False):
|
||||
continue
|
||||
active_names.append(apply_lora_syntax_format(str(lora.get("name") or "")))
|
||||
try:
|
||||
missing = _find_missing_loras(active_names)
|
||||
except Exception:
|
||||
logger.exception("Failed to validate LoRA entries against the local library")
|
||||
return None
|
||||
if not missing:
|
||||
return None
|
||||
return "Missing LoRA(s) in local library: " + ", ".join(missing)
|
||||
|
||||
|
||||
def load_state_dict_in_safetensors(path, device="cpu", filter_prefix=""):
|
||||
"""Simplified version of load_state_dict_in_safetensors that just loads from a local path"""
|
||||
import safetensors.torch
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
from ..utils.utils import get_lora_info_absolute
|
||||
from ..config import config
|
||||
from .utils import FlexibleOptionalInputType, any_type, get_loras_list
|
||||
from .utils import FlexibleOptionalInputType, any_type, get_loras_list, validate_lora_entries
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -31,15 +31,21 @@ class WanVideoLoraSelectLM:
|
||||
"placeholder": "Search LoRAs to add...",
|
||||
"tooltip": "Format: <lora:lora_name:strength> separated by spaces or punctuation",
|
||||
}),
|
||||
"loras": ("LORAS", {}),
|
||||
},
|
||||
"optional": FlexibleOptionalInputType(any_type),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, loras=None):
|
||||
"""Queue-time validation: reject missing local LoRAs before execution."""
|
||||
return validate_lora_entries({"loras": loras}) or True
|
||||
|
||||
RETURN_TYPES = ("WANVIDLORA", "STRING", "STRING")
|
||||
RETURN_NAMES = ("lora", "trigger_words", "active_loras")
|
||||
FUNCTION = "process_loras"
|
||||
|
||||
def process_loras(self, text, low_mem_load=False, merge_loras=True, **kwargs):
|
||||
def process_loras(self, text, loras, low_mem_load=False, merge_loras=True, **kwargs):
|
||||
loras_list = []
|
||||
all_trigger_words = []
|
||||
active_loras = []
|
||||
@@ -57,8 +63,8 @@ class WanVideoLoraSelectLM:
|
||||
selected_blocks = blocks.get("selected_blocks", {})
|
||||
layer_filter = blocks.get("layer_filter", "")
|
||||
|
||||
# Process loras from kwargs with support for both old and new formats
|
||||
loras_from_widget = get_loras_list(kwargs)
|
||||
# Process loras from the widget with support for both old and new formats
|
||||
loras_from_widget = get_loras_list({"loras": loras})
|
||||
for lora in loras_from_widget:
|
||||
if not lora.get('active', False):
|
||||
continue
|
||||
|
||||
+64
-8
@@ -1,3 +1,7 @@
|
||||
# pyright: reportImportCycles=false
|
||||
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||
# import cycles. Breaking them would require an architectural refactor.
|
||||
"""Base classes for recipe parsers."""
|
||||
|
||||
import json
|
||||
@@ -7,7 +11,7 @@ import re
|
||||
from typing import Dict, List, Any, Optional, Tuple
|
||||
from abc import ABC, abstractmethod
|
||||
from ..config import config
|
||||
from ..utils.constants import VALID_LORA_TYPES, VALID_CHECKPOINT_SUB_TYPES
|
||||
from ..utils.constants import MODEL_WEIGHT_FILE_TYPES, VALID_LORA_TYPES, VALID_CHECKPOINT_SUB_TYPES
|
||||
from ..utils.civitai_utils import rewrite_preview_url
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -38,7 +42,41 @@ class RecipeMetadataParser(ABC):
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
async def populate_lora_from_civitai(lora_entry: Dict[str, Any], civitai_info_tuple: Tuple[Dict[str, Any], Optional[str]],
|
||||
def populate_lora_from_local(lora_entry: Dict[str, Any], local_lora: Dict[str, Any], base_model_counts=None) -> Dict[str, Any]:
|
||||
"""Populate a recipe LoRA entry from the local scanner cache."""
|
||||
local_path = local_lora.get('file_path') or ''
|
||||
file_name = local_lora.get('file_name') or os.path.splitext(os.path.basename(local_path))[0]
|
||||
base_model = local_lora.get('base_model') or ''
|
||||
|
||||
lora_entry['name'] = local_lora.get('model_name') or file_name or lora_entry.get('name', '')
|
||||
lora_entry['file_name'] = file_name
|
||||
lora_entry['hash'] = (local_lora.get('sha256') or lora_entry.get('hash') or '').lower()
|
||||
lora_entry['localPath'] = local_path or None
|
||||
lora_entry['size'] = local_lora.get('size', 0) or 0
|
||||
lora_entry['baseModel'] = base_model
|
||||
lora_entry['existsLocally'] = True
|
||||
lora_entry['isDeleted'] = False
|
||||
|
||||
preview_url = local_lora.get('preview_url')
|
||||
if preview_url:
|
||||
lora_entry['thumbnailUrl'] = config.get_preview_static_url(preview_url)
|
||||
|
||||
civitai_info = local_lora.get('civitai') or {}
|
||||
if isinstance(civitai_info, dict):
|
||||
if civitai_info.get('id') is not None:
|
||||
lora_entry['id'] = civitai_info['id']
|
||||
if civitai_info.get('modelId') is not None:
|
||||
lora_entry['modelId'] = civitai_info['modelId']
|
||||
if civitai_info.get('name'):
|
||||
lora_entry['version'] = civitai_info['name']
|
||||
|
||||
if base_model_counts is not None and base_model:
|
||||
base_model_counts[base_model] = base_model_counts.get(base_model, 0) + 1
|
||||
|
||||
return lora_entry
|
||||
|
||||
@staticmethod
|
||||
async def populate_lora_from_civitai(lora_entry: Dict[str, Any], civitai_info_tuple: Tuple[Dict[str, Any] | None, str | None] | Dict[str, Any],
|
||||
recipe_scanner=None, base_model_counts=None, hash_value=None) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Populate a lora entry with information from Civitai API response
|
||||
@@ -151,9 +189,9 @@ class RecipeMetadataParser(ABC):
|
||||
|
||||
# Process file information if available
|
||||
if 'files' in civitai_info:
|
||||
# Find the primary model file (type="Model" and primary=true) in the files list
|
||||
# Find the primary model file (weights-type and primary=true) in the files list
|
||||
model_file = next((file for file in civitai_info.get('files', [])
|
||||
if file.get('type') == 'Model' and file.get('primary') == True), None)
|
||||
if file.get('type') in MODEL_WEIGHT_FILE_TYPES and file.get('primary') == True), None)
|
||||
|
||||
if model_file:
|
||||
# Get size
|
||||
@@ -175,10 +213,18 @@ class RecipeMetadataParser(ABC):
|
||||
lora_entry['localPath'] = local_path
|
||||
lora_entry['file_name'] = os.path.splitext(os.path.basename(local_path))[0]
|
||||
|
||||
# Get thumbnail from local preview if available
|
||||
# Get thumbnail from local preview if available.
|
||||
# Match the cache item by local path first (get_path_by_hash
|
||||
# cascade: 10-char autov2 / 12-char autov3), then by hash.
|
||||
lora_cache = await lora_scanner.get_cached_data()
|
||||
h = (lora_entry.get("hash") or "").lower()
|
||||
lora_item = next((item for item in lora_cache.raw_data
|
||||
if item['sha256'].lower() == lora_entry['hash'].lower()), None)
|
||||
if (item.get("file_path") or "") == local_path), None)
|
||||
if lora_item is None:
|
||||
lora_item = next((item for item in lora_cache.raw_data
|
||||
if (item.get("sha256") or "").lower() == h
|
||||
or (item.get("autov3") or "").lower() == h
|
||||
or (item.get("sha256") or "")[:10].lower() == h), None)
|
||||
if lora_item and 'preview_url' in lora_item:
|
||||
lora_entry['thumbnailUrl'] = config.get_preview_static_url(lora_item['preview_url'])
|
||||
except Exception as e:
|
||||
@@ -194,7 +240,7 @@ class RecipeMetadataParser(ABC):
|
||||
return lora_entry
|
||||
|
||||
@staticmethod
|
||||
async def populate_checkpoint_from_civitai(checkpoint: Dict[str, Any], civitai_info: Dict[str, Any]) -> Dict[str, Any]:
|
||||
async def populate_checkpoint_from_civitai(checkpoint: Dict[str, Any], civitai_info: Dict[str, Any] | Tuple[Dict[str, Any] | None, str | None] | None) -> Dict[str, Any]:
|
||||
"""
|
||||
Populate checkpoint information from Civitai API response
|
||||
|
||||
@@ -249,11 +295,21 @@ class RecipeMetadataParser(ABC):
|
||||
checkpoint['id'] = civitai_data.get('id', 0)
|
||||
|
||||
if 'files' in civitai_data:
|
||||
# Prefer the file CivitAI marked primary; fall back to any
|
||||
# weights-type file (providers without primary flags).
|
||||
model_file = next(
|
||||
(
|
||||
file
|
||||
for file in civitai_data.get('files', [])
|
||||
if file.get('type') == 'Model'
|
||||
if file.get('type') in MODEL_WEIGHT_FILE_TYPES
|
||||
and file.get('primary') is True
|
||||
),
|
||||
None,
|
||||
) or next(
|
||||
(
|
||||
file
|
||||
for file in civitai_data.get('files', [])
|
||||
if file.get('type') in MODEL_WEIGHT_FILE_TYPES
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
@@ -1,3 +1,7 @@
|
||||
# pyright: reportImportCycles=false
|
||||
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||
# import cycles. Breaking them would require an architectural refactor.
|
||||
import logging
|
||||
import json
|
||||
import os
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
"""Factory for creating recipe metadata parsers."""
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
from .parsers import (
|
||||
RecipeFormatParser,
|
||||
ComfyMetadataParser,
|
||||
@@ -31,7 +32,8 @@ class RecipeParserFactory:
|
||||
# First, try CivitaiApiMetadataParser for dict input
|
||||
if isinstance(metadata, dict):
|
||||
try:
|
||||
if CivitaiApiMetadataParser().is_metadata_matching(metadata):
|
||||
user_comment: Any = metadata
|
||||
if CivitaiApiMetadataParser().is_metadata_matching(user_comment):
|
||||
return CivitaiApiMetadataParser()
|
||||
except Exception as e:
|
||||
logger.debug(f"CivitaiApiMetadataParser check failed: {e}")
|
||||
|
||||
+197
-38
@@ -8,6 +8,7 @@ from typing import Dict, Any
|
||||
from ..base import RecipeMetadataParser
|
||||
from ..constants import GEN_PARAM_KEYS
|
||||
from ...services.metadata_service import get_default_metadata_provider
|
||||
from ...utils.constants import is_empty_placeholder_hash
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -52,7 +53,7 @@ class AutomaticMetadataParser(RecipeMetadataParser):
|
||||
negative_and_params = ""
|
||||
|
||||
# Initialize metadata
|
||||
metadata = {
|
||||
metadata: Dict[str, Any] = {
|
||||
"prompt": prompt,
|
||||
"loras": []
|
||||
}
|
||||
@@ -146,14 +147,12 @@ class AutomaticMetadataParser(RecipeMetadataParser):
|
||||
# Initialize hashes dict if it doesn't exist
|
||||
if "hashes" not in metadata:
|
||||
metadata["hashes"] = {}
|
||||
# Add as lora type in the same format as
|
||||
# regular hashes. Only override an
|
||||
# existing entry if its value is empty
|
||||
# (Lora hashes is the more reliable
|
||||
# source when Hashes JSON has blanks).
|
||||
# Lora hashes carries the 12-char AutoV3
|
||||
# hash (resolvable on CivitAI and the local
|
||||
# autov3 index); the Hashes JSON value is
|
||||
# only the 10-char AutoV2 prefix, so on
|
||||
# conflict the Lora hashes value wins.
|
||||
key = f"lora:{lora_name}"
|
||||
existing = metadata["hashes"].get(key, "")
|
||||
if not existing:
|
||||
metadata["hashes"][key] = lora_hash
|
||||
|
||||
# Remove lora hashes from params section
|
||||
@@ -362,37 +361,47 @@ class AutomaticMetadataParser(RecipeMetadataParser):
|
||||
|
||||
checkpoint = checkpoint_entry
|
||||
|
||||
# If no LoRAs from Civitai resources or to supplement, extract from metadata["hashes"]
|
||||
if not loras or len(loras) == 0:
|
||||
# Extract lora weights from extranet tags in prompt (for later use)
|
||||
lora_weights = {}
|
||||
lora_matches = re.findall(self.EXTRANETS_REGEX, prompt)
|
||||
for lora_type, lora_name, lora_weight in lora_matches:
|
||||
key = f"{lora_type}:{lora_name}"
|
||||
lora_weights[key] = round(float(lora_weight), 2)
|
||||
def normalize_lora_name(name, basename=False):
|
||||
normalized = str(name or '').replace('\\', '/')
|
||||
if normalized.casefold().endswith('.safetensors'):
|
||||
normalized = normalized[:-12]
|
||||
if basename:
|
||||
normalized = normalized.rsplit('/', 1)[-1]
|
||||
return normalized.casefold()
|
||||
|
||||
# Use hashes from metadata as the primary source
|
||||
if metadata.get("hashes"):
|
||||
for hash_key, lora_hash in metadata.get("hashes", {}).items():
|
||||
# Only process lora or hypernet types
|
||||
if not hash_key.startswith(("lora:", "hypernet:")):
|
||||
continue
|
||||
def get_version_id(lora):
|
||||
version_id = lora.get('id')
|
||||
if version_id in (None, '', 0, '0'):
|
||||
version_id = lora.get('modelVersionId')
|
||||
if version_id in (None, '', 0, '0'):
|
||||
return None
|
||||
return str(version_id)
|
||||
|
||||
# Skip entries without a hash value — they can't be
|
||||
# resolved via CivitAI and would only produce a
|
||||
# useless "Deleted" entry in the recipe.
|
||||
if not lora_hash:
|
||||
continue
|
||||
prompt_loras = {}
|
||||
for match in re.findall(self.EXTRANETS_REGEX, prompt):
|
||||
lora_type, lora_name, _ = match
|
||||
prompt_loras[(lora_type, normalize_lora_name(lora_name))] = match
|
||||
|
||||
lora_type, lora_name = hash_key.split(':', 1)
|
||||
prompt_by_basename = {}
|
||||
for lora_type, lora_name, lora_weight in prompt_loras.values():
|
||||
key = (lora_type, normalize_lora_name(lora_name, True))
|
||||
prompt_by_basename.setdefault(key, []).append((lora_name, round(float(lora_weight), 2)))
|
||||
|
||||
# Get weight from extranet tags if available, else default to 1.0
|
||||
weight = lora_weights.get(hash_key, 1.0)
|
||||
hash_basenames = {
|
||||
(hash_key.split(':', 1)[0], normalize_lora_name(hash_key.split(':', 1)[1], True))
|
||||
for hash_key, hash_value in metadata.get("hashes", {}).items()
|
||||
if hash_value and hash_key.startswith(("lora:", "hypernet:"))
|
||||
}
|
||||
recipe_base_model = checkpoint.get("baseModel") if checkpoint else None
|
||||
if not recipe_base_model and len(base_model_counts) == 1:
|
||||
recipe_base_model = next(iter(base_model_counts))
|
||||
|
||||
# Initialize lora entry
|
||||
lora_entry = {
|
||||
resource_lora_count = len(loras)
|
||||
|
||||
def make_lora_entry(lora_type, lora_name, weight, lora_hash=''):
|
||||
return {
|
||||
'name': lora_name,
|
||||
'type': lora_type, # 'lora' or 'hypernet'
|
||||
'type': lora_type,
|
||||
'weight': weight,
|
||||
'hash': lora_hash,
|
||||
'existsLocally': False,
|
||||
@@ -405,24 +414,174 @@ class AutomaticMetadataParser(RecipeMetadataParser):
|
||||
'isDeleted': False
|
||||
}
|
||||
|
||||
# Try to get info from Civitai
|
||||
if metadata_provider:
|
||||
def merge_or_append_civitai(civitai_entry, preserve_existing_weight=False):
|
||||
civitai_id = get_version_id(civitai_entry)
|
||||
civitai_hash = (civitai_entry.get('hash') or '').lower()
|
||||
for index, existing in enumerate(loras):
|
||||
existing_id = get_version_id(existing)
|
||||
existing_hash = (existing.get('hash') or '').lower()
|
||||
if not (
|
||||
(civitai_id and existing_id == civitai_id)
|
||||
or (civitai_hash and existing_hash == civitai_hash)
|
||||
):
|
||||
continue
|
||||
|
||||
if preserve_existing_weight:
|
||||
civitai_entry['weight'] = existing.get('weight', civitai_entry['weight'])
|
||||
existing_base = existing.get('baseModel')
|
||||
if not civitai_entry.get('baseModel'):
|
||||
civitai_entry['baseModel'] = existing_base or ''
|
||||
elif existing_base:
|
||||
remaining = base_model_counts.get(existing_base, 0) - 1
|
||||
if remaining > 0:
|
||||
base_model_counts[existing_base] = remaining
|
||||
else:
|
||||
base_model_counts.pop(existing_base, None)
|
||||
loras[index] = civitai_entry
|
||||
return
|
||||
loras.append(civitai_entry)
|
||||
|
||||
def merge_or_append_local(local_entry):
|
||||
local_id = get_version_id(local_entry)
|
||||
local_hash = (local_entry.get('hash') or '').lower()
|
||||
for existing in loras:
|
||||
existing_id = get_version_id(existing)
|
||||
existing_hash = (existing.get('hash') or '').lower()
|
||||
if not (
|
||||
(local_id and existing_id == local_id)
|
||||
or (local_hash and existing_hash == local_hash)
|
||||
):
|
||||
continue
|
||||
|
||||
existing['weight'] = local_entry['weight']
|
||||
existing['hash'] = local_entry['hash']
|
||||
existing['file_name'] = local_entry['file_name']
|
||||
existing['existsLocally'] = True
|
||||
existing['localPath'] = local_entry['localPath']
|
||||
existing['size'] = local_entry['size']
|
||||
existing['isDeleted'] = False
|
||||
if not existing.get('modelId') and local_entry.get('modelId'):
|
||||
existing['modelId'] = local_entry['modelId']
|
||||
if not existing.get('baseModel') and local_entry.get('baseModel'):
|
||||
existing['baseModel'] = local_entry['baseModel']
|
||||
base_model_counts[local_entry['baseModel']] = base_model_counts.get(local_entry['baseModel'], 0) + 1
|
||||
thumbnail_url = local_entry.get('thumbnailUrl')
|
||||
if thumbnail_url and not thumbnail_url.endswith('/images/no-preview.png'):
|
||||
existing['thumbnailUrl'] = thumbnail_url
|
||||
return
|
||||
|
||||
if local_entry.get('baseModel'):
|
||||
base_model = local_entry['baseModel']
|
||||
base_model_counts[base_model] = base_model_counts.get(base_model, 0) + 1
|
||||
loras.append(local_entry)
|
||||
|
||||
resolved_prompt_basenames = set()
|
||||
queried_local_basenames = set()
|
||||
for lora_type, lora_name, lora_weight in prompt_loras.values():
|
||||
weight = round(float(lora_weight), 2)
|
||||
basename_key = (lora_type, normalize_lora_name(lora_name, True))
|
||||
matching_resources = [
|
||||
lora
|
||||
for lora in loras[:resource_lora_count]
|
||||
if lora.get('file_name')
|
||||
and normalize_lora_name(lora['file_name'], True) == basename_key[1]
|
||||
and (
|
||||
(lora_type == 'hypernet' and str(lora.get('type', '')).casefold() in ('hypernet', 'hypernetwork'))
|
||||
or (lora_type == 'lora' and str(lora.get('type', '')).casefold() not in ('hypernet', 'hypernetwork'))
|
||||
)
|
||||
]
|
||||
if len(prompt_by_basename[basename_key]) == 1 and len(matching_resources) == 1:
|
||||
matching_resources[0]['weight'] = weight
|
||||
if basename_key not in hash_basenames:
|
||||
resolved_prompt_basenames.add(basename_key)
|
||||
continue
|
||||
|
||||
if basename_key in hash_basenames:
|
||||
continue
|
||||
|
||||
if not recipe_scanner or lora_type != 'lora':
|
||||
continue
|
||||
queried_local_basenames.add(basename_key)
|
||||
local_lora = await recipe_scanner.get_local_lora(lora_name, recipe_base_model)
|
||||
if not local_lora:
|
||||
continue
|
||||
|
||||
local_entry = self.populate_lora_from_local(
|
||||
make_lora_entry(lora_type, lora_name, weight),
|
||||
local_lora,
|
||||
)
|
||||
merge_or_append_local(local_entry)
|
||||
resolved_prompt_basenames.add(basename_key)
|
||||
|
||||
for hash_key, lora_hash in metadata.get("hashes", {}).items():
|
||||
if not hash_key.startswith(("lora:", "hypernet:")):
|
||||
continue
|
||||
lora_type, lora_name = hash_key.split(':', 1)
|
||||
basename_key = (lora_type, normalize_lora_name(lora_name, True))
|
||||
if basename_key in resolved_prompt_basenames:
|
||||
continue
|
||||
|
||||
prompt_entries = prompt_by_basename.get(basename_key, [])
|
||||
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)
|
||||
|
||||
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':
|
||||
local_lora = await recipe_scanner.get_local_lora_by_hash(lora_hash)
|
||||
if local_lora:
|
||||
local_entry = self.populate_lora_from_local(lora_entry, local_lora)
|
||||
merge_or_append_local(local_entry)
|
||||
continue
|
||||
|
||||
hash_resolved = False
|
||||
if lora_hash and metadata_provider:
|
||||
try:
|
||||
civitai_info = await metadata_provider.get_model_by_hash(lora_hash)
|
||||
|
||||
populated_entry = await self.populate_lora_from_civitai(
|
||||
lora_entry,
|
||||
civitai_info,
|
||||
recipe_scanner,
|
||||
base_model_counts,
|
||||
lora_hash
|
||||
lora_hash,
|
||||
)
|
||||
if populated_entry is None:
|
||||
continue # Skip invalid LoRA types
|
||||
continue
|
||||
lora_entry = populated_entry
|
||||
hash_resolved = not lora_entry.get('isDeleted')
|
||||
except Exception as e:
|
||||
logger.error(f"Error fetching Civitai info for LoRA {lora_name}: {e}")
|
||||
|
||||
if hash_resolved:
|
||||
merge_or_append_civitai(lora_entry, preserve_existing_weight=not prompt_entries)
|
||||
continue
|
||||
|
||||
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
|
||||
|
||||
if lora_hash and not resource_lora_count:
|
||||
loras.append(lora_entry)
|
||||
|
||||
# Try to get base model from resources or make educated guess
|
||||
|
||||
@@ -4,7 +4,7 @@ import json
|
||||
import logging
|
||||
from typing import Dict, Any, Union
|
||||
from ..base import RecipeMetadataParser
|
||||
from ..constants import GEN_PARAM_KEYS
|
||||
from ..constants import GEN_PARAM_KEYS, VALID_LORA_TYPES
|
||||
from ...services.metadata_service import get_default_metadata_provider
|
||||
from ...config import config
|
||||
|
||||
@@ -14,15 +14,16 @@ logger = logging.getLogger(__name__)
|
||||
class CivitaiApiMetadataParser(RecipeMetadataParser):
|
||||
"""Parser for Civitai image metadata format"""
|
||||
|
||||
def is_metadata_matching(self, metadata) -> bool:
|
||||
def is_metadata_matching(self, user_comment) -> bool:
|
||||
"""Check if the metadata matches the Civitai image metadata format
|
||||
|
||||
Args:
|
||||
metadata: The metadata from the image (dict)
|
||||
user_comment: The metadata from the image (dict)
|
||||
|
||||
Returns:
|
||||
bool: True if this parser can handle the metadata
|
||||
"""
|
||||
metadata = user_comment
|
||||
if not metadata or not isinstance(metadata, dict):
|
||||
return False
|
||||
|
||||
@@ -73,7 +74,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
||||
|
||||
return False
|
||||
|
||||
async def parse_metadata( # type: ignore[override]
|
||||
async def parse_metadata( # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
self, user_comment, recipe_scanner=None, civitai_client=None,
|
||||
local_cache: dict[str, Any] | None = None,
|
||||
) -> Dict[str, Any]:
|
||||
@@ -89,8 +90,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
||||
Returns:
|
||||
Dict containing parsed recipe data
|
||||
"""
|
||||
metadata: Dict[str, Any] = user_comment # type: ignore[assignment]
|
||||
metadata = user_comment
|
||||
metadata: Dict[str, Any] = user_comment
|
||||
try:
|
||||
# Get metadata provider instead of using civitai_client directly
|
||||
metadata_provider = await get_default_metadata_provider()
|
||||
@@ -115,8 +115,29 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
||||
):
|
||||
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
|
||||
result = {
|
||||
result: Dict[str, Any] = {
|
||||
"base_model": None,
|
||||
"loras": [],
|
||||
"model": None,
|
||||
@@ -125,10 +146,10 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
||||
}
|
||||
|
||||
# Track already added LoRAs to prevent duplicates
|
||||
added_loras = {} # key: model_version_id or hash, value: index in result["loras"]
|
||||
added_loras: Dict[str, Any] = {} # key: model_version_id or hash, value: index in result["loras"]
|
||||
|
||||
# Extract hash information from hashes field for LoRA matching
|
||||
lora_hashes = {}
|
||||
lora_hashes: Dict[str, Any] = {}
|
||||
if "hashes" in metadata and isinstance(metadata["hashes"], dict):
|
||||
for key, hash_value in metadata["hashes"].items():
|
||||
key_str = str(key)
|
||||
@@ -184,7 +205,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
||||
if model_info:
|
||||
result["base_model"] = model_info.get("baseModel", "")
|
||||
|
||||
base_model_counts = {}
|
||||
base_model_counts: Dict[str, int] = {}
|
||||
|
||||
# Process standard resources array
|
||||
if "resources" in metadata and isinstance(metadata["resources"], list):
|
||||
@@ -196,7 +217,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
||||
# identification because it has an explicit type field and hash,
|
||||
# unlike modelVersionIds which is a flat list with no type info.
|
||||
if resource_type == "model":
|
||||
checkpoint_entry = {
|
||||
checkpoint_entry: Dict[str, Any] = {
|
||||
"id": 0,
|
||||
"modelId": 0,
|
||||
"name": resource.get("name", "Unknown Model"),
|
||||
@@ -216,7 +237,8 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
||||
# Try to look up base model from the checkpoint hash
|
||||
cp_hash = checkpoint_entry.get("hash")
|
||||
if cp_hash and metadata_provider:
|
||||
local_cached = local_cache.get(cp_hash) if local_cache else None
|
||||
# local_cache keys are stored lowercase
|
||||
local_cached = local_cache.get(cp_hash.lower()) if local_cache else None
|
||||
if local_cached:
|
||||
self._populate_entry_from_cache(
|
||||
checkpoint_entry, local_cached
|
||||
@@ -294,8 +316,15 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
||||
|
||||
# Try to get info from Civitai if hash is available
|
||||
if lora_hash and metadata_provider:
|
||||
local_cached = local_cache.get(lora_hash) if local_cache else None
|
||||
# local_cache keys are stored lowercase
|
||||
local_cached = local_cache.get(lora_hash.lower()) if local_cache else None
|
||||
if local_cached:
|
||||
cached_type = self._cache_item_model_type(local_cached)
|
||||
if cached_type and cached_type not in VALID_LORA_TYPES:
|
||||
logger.debug(
|
||||
f"Skipping non-LoRA cache item for hash {lora_hash}"
|
||||
)
|
||||
continue
|
||||
self._populate_entry_from_cache(
|
||||
lora_entry, local_cached
|
||||
)
|
||||
@@ -304,6 +333,12 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
||||
added_loras[str(lora_entry["id"])] = len(
|
||||
result["loras"]
|
||||
)
|
||||
# Mirror base.py:150-151 counts for API-path loras
|
||||
bm = local_cached.get("base_model") or ""
|
||||
if bm:
|
||||
base_model_counts[bm] = base_model_counts.get(
|
||||
bm, 0
|
||||
) + 1
|
||||
else:
|
||||
try:
|
||||
civitai_info = (
|
||||
@@ -649,6 +684,23 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
||||
}
|
||||
|
||||
if metadata_provider:
|
||||
# local_cache keys are stored lowercase
|
||||
local_cached = local_cache.get(lora_hash.lower()) if local_cache else None
|
||||
if local_cached:
|
||||
cached_type = self._cache_item_model_type(local_cached)
|
||||
if cached_type and cached_type not in VALID_LORA_TYPES:
|
||||
logger.debug(
|
||||
f"Skipping non-LoRA cache item for hash {lora_hash}"
|
||||
)
|
||||
continue
|
||||
self._populate_entry_from_cache(lora_entry, local_cached)
|
||||
# Mirror base.py:150-151 counts for API-path loras
|
||||
bm = local_cached.get("base_model") or ""
|
||||
if bm:
|
||||
base_model_counts[bm] = base_model_counts.get(bm, 0) + 1
|
||||
if "id" in lora_entry and lora_entry["id"]:
|
||||
added_loras[str(lora_entry["id"])] = len(result["loras"])
|
||||
else:
|
||||
try:
|
||||
civitai_info = await metadata_provider.get_model_by_hash(
|
||||
lora_hash
|
||||
@@ -711,6 +763,25 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
||||
|
||||
# Try to get info from Civitai if hash is available
|
||||
if lora_entry["hash"] and metadata_provider:
|
||||
# local_cache keys are stored lowercase
|
||||
local_cached = local_cache.get(lora_hash.lower()) if local_cache else None
|
||||
if local_cached:
|
||||
cached_type = self._cache_item_model_type(local_cached)
|
||||
if cached_type and cached_type not in VALID_LORA_TYPES:
|
||||
logger.debug(
|
||||
f"Skipping non-LoRA cache item for hash {lora_hash}"
|
||||
)
|
||||
lora_index += 1
|
||||
continue # Skip non-LoRA cache items
|
||||
self._populate_entry_from_cache(lora_entry, local_cached)
|
||||
# Mirror base.py:150-151 counts for API-path loras
|
||||
bm = local_cached.get("base_model") or ""
|
||||
if bm:
|
||||
base_model_counts[bm] = base_model_counts.get(bm, 0) + 1
|
||||
# If we have a version ID from Civitai, track it for deduplication
|
||||
if "id" in lora_entry and lora_entry["id"]:
|
||||
added_loras[str(lora_entry["id"])] = len(result["loras"])
|
||||
else:
|
||||
try:
|
||||
civitai_info = await metadata_provider.get_model_by_hash(
|
||||
lora_hash
|
||||
@@ -795,3 +866,14 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
||||
base_model = cache_item.get("base_model", "")
|
||||
if base_model:
|
||||
entry["baseModel"] = base_model
|
||||
|
||||
@staticmethod
|
||||
def _cache_item_model_type(cache_item: dict[str, Any]) -> str:
|
||||
"""Lowercased civitai.model.type of a cache item, or '' when unknown."""
|
||||
civ = cache_item.get("civitai")
|
||||
if not isinstance(civ, dict):
|
||||
return ""
|
||||
model_info = civ.get("model")
|
||||
if not isinstance(model_info, dict):
|
||||
return ""
|
||||
return (model_info.get("type") or "").lower()
|
||||
|
||||
+104
-67
@@ -31,79 +31,25 @@ class ComfyMetadataParser(RecipeMetadataParser):
|
||||
metadata_provider = await get_default_metadata_provider()
|
||||
|
||||
data = json.loads(user_comment)
|
||||
loras = []
|
||||
|
||||
# Find all LoraLoader nodes
|
||||
lora_nodes = {k: v for k, v in data.items() if isinstance(v, dict) and v.get('class_type') == 'LoraLoader'}
|
||||
|
||||
# Process each LoraLoader node
|
||||
for node_id, node in lora_nodes.items():
|
||||
if 'inputs' not in node or 'lora_name' not in node['inputs']:
|
||||
continue
|
||||
|
||||
lora_name = node['inputs'].get('lora_name', '')
|
||||
|
||||
# Parse the URN to extract model ID and version ID
|
||||
# Format: "urn:air:sdxl:lora:civitai:1107767@1253442"
|
||||
lora_id_match = re.search(r'civitai:(\d+)@(\d+)', lora_name)
|
||||
if not lora_id_match:
|
||||
continue
|
||||
|
||||
model_id = lora_id_match.group(1)
|
||||
model_version_id = lora_id_match.group(2)
|
||||
|
||||
# Get strength from node inputs
|
||||
weight = node['inputs'].get('strength_model', 1.0)
|
||||
|
||||
# Initialize lora entry with default values
|
||||
lora_entry = {
|
||||
'id': model_version_id,
|
||||
'modelId': model_id,
|
||||
'name': f"Lora {model_id}", # Default name
|
||||
'version': '',
|
||||
'type': 'lora',
|
||||
'weight': weight,
|
||||
'existsLocally': False,
|
||||
'localPath': None,
|
||||
'file_name': '',
|
||||
'hash': '',
|
||||
'thumbnailUrl': '/loras_static/images/no-preview.png',
|
||||
'baseModel': '',
|
||||
'size': 0,
|
||||
'downloadUrl': '',
|
||||
'isDeleted': False
|
||||
}
|
||||
|
||||
# Get additional info from Civitai if metadata provider is available
|
||||
if metadata_provider:
|
||||
try:
|
||||
civitai_info_tuple = await metadata_provider.get_model_version_info(model_version_id)
|
||||
# Populate lora entry with Civitai info
|
||||
populated_entry = await self.populate_lora_from_civitai(
|
||||
lora_entry,
|
||||
civitai_info_tuple,
|
||||
recipe_scanner
|
||||
)
|
||||
if populated_entry is None:
|
||||
continue # Skip invalid LoRA types
|
||||
lora_entry = populated_entry
|
||||
except Exception as e:
|
||||
logger.error(f"Error fetching Civitai info for LoRA: {e}")
|
||||
|
||||
loras.append(lora_entry)
|
||||
|
||||
# Find checkpoint info
|
||||
checkpoint_nodes = {k: v for k, v in data.items() if isinstance(v, dict) and v.get('class_type') == 'CheckpointLoaderSimple'}
|
||||
checkpoint = None
|
||||
checkpoint_id = None
|
||||
checkpoint_version_id = None
|
||||
|
||||
if checkpoint_nodes:
|
||||
# Get the first checkpoint node
|
||||
checkpoint_node = next(iter(checkpoint_nodes.values()))
|
||||
if 'inputs' in checkpoint_node and 'ckpt_name' in checkpoint_node['inputs']:
|
||||
checkpoint_name = checkpoint_node['inputs']['ckpt_name']
|
||||
# Parse checkpoint URN
|
||||
# Some ComfyUI workflows serialize ckpt_name as a
|
||||
# single-element list (e.g. ["model.safetensors"]) or leave
|
||||
# the value unset (None). Neither is a string, so skip the
|
||||
# CivitAI-URN lookup instead of crashing re.search with a
|
||||
# TypeError that fails the whole image import.
|
||||
if isinstance(checkpoint_name, list):
|
||||
checkpoint_name = (
|
||||
checkpoint_name[0] if checkpoint_name else None
|
||||
)
|
||||
if isinstance(checkpoint_name, str):
|
||||
checkpoint_match = re.search(r'civitai:(\d+)@(\d+)', checkpoint_name)
|
||||
if checkpoint_match:
|
||||
checkpoint_id = checkpoint_match.group(1)
|
||||
@@ -115,17 +61,108 @@ class ComfyMetadataParser(RecipeMetadataParser):
|
||||
'version': '',
|
||||
'type': 'checkpoint'
|
||||
}
|
||||
|
||||
# Get additional checkpoint info from Civitai
|
||||
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)
|
||||
# Populate checkpoint with Civitai info
|
||||
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
|
||||
loras = []
|
||||
lora_candidates = []
|
||||
for node in data.values():
|
||||
if not isinstance(node, dict):
|
||||
continue
|
||||
|
||||
inputs = node.get('inputs')
|
||||
if not isinstance(inputs, dict):
|
||||
continue
|
||||
|
||||
if node.get('class_type') == 'LoraLoader':
|
||||
lora_name = inputs.get('lora_name', '')
|
||||
if isinstance(lora_name, str) and lora_name:
|
||||
lora_candidates.append((lora_name, inputs.get('strength_model', 1.0)))
|
||||
continue
|
||||
|
||||
if node.get('class_type') != 'LoraLoaderLM':
|
||||
continue
|
||||
|
||||
loras_data = inputs.get('loras', [])
|
||||
if isinstance(loras_data, dict):
|
||||
loras_data = loras_data.get('__value__', [])
|
||||
if isinstance(loras_data, list) and len(loras_data) == 1 and isinstance(loras_data[0], list):
|
||||
loras_data = loras_data[0]
|
||||
if not isinstance(loras_data, list):
|
||||
continue
|
||||
|
||||
for lora in loras_data:
|
||||
if not isinstance(lora, dict) or not lora.get('active', False) or lora.get('_isDummy', False):
|
||||
continue
|
||||
lora_name = lora.get('name', '')
|
||||
if isinstance(lora_name, str) and lora_name:
|
||||
lora_candidates.append((lora_name, lora.get('strength', 1.0)))
|
||||
|
||||
for lora_name, weight in lora_candidates:
|
||||
if isinstance(weight, str):
|
||||
try:
|
||||
weight = float(weight)
|
||||
except ValueError:
|
||||
weight = 1.0
|
||||
lora_id_match = re.search(r'civitai:(\d+)@(\d+)', lora_name)
|
||||
if lora_id_match:
|
||||
model_id = lora_id_match.group(1)
|
||||
model_version_id = lora_id_match.group(2)
|
||||
entry_name = f"Lora {model_id}"
|
||||
else:
|
||||
model_id = 0
|
||||
model_version_id = 0
|
||||
entry_name = re.split(r'[\\/]', lora_name)[-1]
|
||||
entry_name = re.sub(r'\.[^.]+$', '', entry_name)
|
||||
|
||||
lora_entry = {
|
||||
'id': model_version_id,
|
||||
'modelId': model_id,
|
||||
'name': entry_name,
|
||||
'version': '',
|
||||
'type': 'lora',
|
||||
'weight': weight,
|
||||
'existsLocally': False,
|
||||
'localPath': None,
|
||||
'file_name': entry_name,
|
||||
'hash': '',
|
||||
'thumbnailUrl': '/loras_static/images/no-preview.png',
|
||||
'baseModel': '',
|
||||
'size': 0,
|
||||
'downloadUrl': '',
|
||||
'isDeleted': False
|
||||
}
|
||||
|
||||
if lora_id_match:
|
||||
if metadata_provider:
|
||||
try:
|
||||
civitai_info_tuple = await metadata_provider.get_model_version_info(model_version_id)
|
||||
populated_entry = await self.populate_lora_from_civitai(
|
||||
lora_entry,
|
||||
civitai_info_tuple,
|
||||
recipe_scanner
|
||||
)
|
||||
if populated_entry is None:
|
||||
continue
|
||||
lora_entry = populated_entry
|
||||
except Exception as e:
|
||||
logger.error(f"Error fetching Civitai info for LoRA: {e}")
|
||||
else:
|
||||
if not recipe_scanner:
|
||||
continue
|
||||
local_lora = await recipe_scanner.get_local_lora(lora_name, recipe_base_model)
|
||||
if not local_lora:
|
||||
continue
|
||||
lora_entry = self.populate_lora_from_local(lora_entry, local_lora)
|
||||
|
||||
loras.append(lora_entry)
|
||||
|
||||
# Extract generation parameters
|
||||
gen_params = {}
|
||||
|
||||
|
||||
@@ -30,7 +30,7 @@ class MetaFormatParser(RecipeMetadataParser):
|
||||
prompt = parts[0].strip()
|
||||
|
||||
# Initialize metadata
|
||||
metadata = {"prompt": prompt, "loras": []}
|
||||
metadata: Dict[str, Any] = {"prompt": prompt, "loras": []}
|
||||
|
||||
# Extract negative prompt and parameters if available
|
||||
if len(parts) > 1:
|
||||
|
||||
@@ -91,7 +91,15 @@ class RecipeFormatParser(RecipeMetadataParser):
|
||||
exists_locally = lora_scanner.has_hash(lora['hash'])
|
||||
if exists_locally:
|
||||
lora_cache = await lora_scanner.get_cached_data()
|
||||
lora_item = next((item for item in lora_cache.raw_data if item['sha256'].lower() == lora['hash'].lower()), None)
|
||||
# Cascade match: full sha256, stored autov3, or autov2 (sha256[:10]).
|
||||
h = (lora.get('hash') or '').lower()
|
||||
lora_item = next(
|
||||
(item for item in lora_cache.raw_data
|
||||
if (item.get("sha256") or "").lower() == h
|
||||
or (item.get("autov3") or "").lower() == h
|
||||
or (item.get("sha256") or "")[:10].lower() == h),
|
||||
None
|
||||
)
|
||||
if lora_item:
|
||||
lora_entry['existsLocally'] = True
|
||||
lora_entry['inLibrary'] = True
|
||||
@@ -148,7 +156,7 @@ class RecipeFormatParser(RecipeMetadataParser):
|
||||
checkpoint_data = recipe_metadata.get('checkpoint') or {}
|
||||
if isinstance(checkpoint_data, dict) and checkpoint_data:
|
||||
version_id = checkpoint_data.get('modelVersionId') or checkpoint_data.get('id')
|
||||
checkpoint_entry = {
|
||||
checkpoint_entry: Dict[str, Any] = {
|
||||
'id': version_id or 0,
|
||||
'modelId': checkpoint_data.get('modelId', 0),
|
||||
'name': checkpoint_data.get('name', 'Unknown Checkpoint'),
|
||||
@@ -188,7 +196,7 @@ class RecipeFormatParser(RecipeMetadataParser):
|
||||
filtered_gen_params[key] = value
|
||||
|
||||
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,
|
||||
'gen_params': filtered_gen_params,
|
||||
'tags': recipe_metadata.get('tags', []),
|
||||
@@ -200,3 +208,24 @@ class RecipeFormatParser(RecipeMetadataParser):
|
||||
except Exception as e:
|
||||
logger.error(f"Error parsing recipe format metadata: {e}", exc_info=True)
|
||||
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()
|
||||
|
||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import TYPE_CHECKING, Callable, Dict, Mapping
|
||||
from typing import TYPE_CHECKING, Awaitable, Callable, Dict, Mapping
|
||||
|
||||
import jinja2
|
||||
from aiohttp import web
|
||||
@@ -30,6 +30,7 @@ from ..services.websocket_progress_callback import (
|
||||
WebSocketProgressCallback,
|
||||
)
|
||||
from ..utils.exif_utils import ExifUtils
|
||||
from ..utils.constants import MODEL_WEIGHT_FILE_TYPES
|
||||
from ..utils.metadata_manager import MetadataManager
|
||||
from .model_route_registrar import COMMON_ROUTE_DEFINITIONS, ModelRouteRegistrar
|
||||
from .handlers.model_handlers import (
|
||||
@@ -84,7 +85,7 @@ class BaseModelRoutes(ABC):
|
||||
self.metadata_progress_callback = WebSocketBroadcastCallback()
|
||||
|
||||
self._handler_set: ModelHandlerSet | None = None
|
||||
self._handler_mapping: Dict[str, Callable[[web.Request], web.StreamResponse]] | None = None
|
||||
self._handler_mapping: Dict[str, Callable[[web.Request], Awaitable[web.Response]]] | None = None
|
||||
|
||||
self._preview_service = PreviewAssetService(
|
||||
metadata_manager=MetadataManager,
|
||||
@@ -131,7 +132,7 @@ class BaseModelRoutes(ABC):
|
||||
self._handler_set = None
|
||||
self._handler_mapping = None
|
||||
|
||||
def _ensure_handler_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
|
||||
def _ensure_handler_mapping(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
|
||||
if self._handler_mapping is None:
|
||||
handler_set = self._create_handler_set()
|
||||
self._handler_set = handler_set
|
||||
@@ -220,7 +221,7 @@ class BaseModelRoutes(ABC):
|
||||
)
|
||||
|
||||
@property
|
||||
def route_handlers(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
|
||||
def route_handlers(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
|
||||
return self._ensure_handler_mapping()
|
||||
|
||||
def setup_routes(self, app: web.Application, prefix: str) -> None:
|
||||
@@ -237,7 +238,7 @@ class BaseModelRoutes(ABC):
|
||||
"""Setup model-specific routes."""
|
||||
raise NotImplementedError
|
||||
|
||||
def _parse_specific_params(self, request: web.Request) -> Dict:
|
||||
def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]:
|
||||
"""Parse model-specific parameters - to be overridden by subclasses."""
|
||||
return {}
|
||||
|
||||
@@ -251,9 +252,9 @@ class BaseModelRoutes(ABC):
|
||||
|
||||
def _find_model_file(self, files):
|
||||
"""Find the appropriate model file from the files list - can be overridden by subclasses."""
|
||||
return next((file for file in files if file.get("type") in ("Model", "Diffusion Model") and file.get("primary") is True), None)
|
||||
return next((file for file in files if file.get("type") in MODEL_WEIGHT_FILE_TYPES and file.get("primary") is True), None)
|
||||
|
||||
def get_handler(self, name: str) -> Callable[[web.Request], web.StreamResponse]:
|
||||
def get_handler(self, name: str) -> Callable[[web.Request], Awaitable[web.StreamResponse]]:
|
||||
"""Expose handlers for subclasses or tests."""
|
||||
return self._ensure_handler_mapping()[name]
|
||||
|
||||
@@ -285,7 +286,7 @@ class BaseModelRoutes(ABC):
|
||||
)
|
||||
return self.model_lifecycle_service
|
||||
|
||||
def _make_handler_proxy(self, name: str) -> Callable[[web.Request], web.StreamResponse]:
|
||||
def _make_handler_proxy(self, name: str) -> Callable[[web.Request], Awaitable[web.StreamResponse]]:
|
||||
async def proxy(request: web.Request) -> web.StreamResponse:
|
||||
try:
|
||||
handler = self.get_handler(name)
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
from typing import Callable, Mapping
|
||||
from typing import Awaitable, Callable, Mapping
|
||||
|
||||
import jinja2
|
||||
from aiohttp import web
|
||||
@@ -32,6 +32,7 @@ from .handlers.recipe_handlers import (
|
||||
RecipePageView,
|
||||
RecipeQueryHandler,
|
||||
RecipeSharingHandler,
|
||||
RecipeWorkflowHandler,
|
||||
)
|
||||
from .recipe_route_registrar import ROUTE_DEFINITIONS
|
||||
|
||||
@@ -61,7 +62,9 @@ class BaseRecipeRoutes:
|
||||
self._i18n_registered = False
|
||||
self._startup_hooks_registered = False
|
||||
self._handler_set: RecipeHandlerSet | None = None
|
||||
self._handler_mapping: dict[str, Callable] | None = None
|
||||
self._handler_mapping: Mapping[
|
||||
str, Callable[[web.Request], Awaitable[web.StreamResponse]]
|
||||
] | None = None
|
||||
|
||||
async def attach_dependencies(self, app: web.Application | None = None) -> None:
|
||||
"""Resolve shared services from the registry."""
|
||||
@@ -84,7 +87,9 @@ class BaseRecipeRoutes:
|
||||
app.on_startup.append(self.attach_dependencies)
|
||||
self._startup_hooks_registered = True
|
||||
|
||||
def to_route_mapping(self) -> Mapping[str, Callable]:
|
||||
def to_route_mapping(
|
||||
self,
|
||||
) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
|
||||
"""Return a mapping of handler name to coroutine for registrar binding."""
|
||||
|
||||
if self._handler_mapping is None:
|
||||
@@ -124,17 +129,17 @@ class BaseRecipeRoutes:
|
||||
or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0"
|
||||
)
|
||||
if not standalone_mode:
|
||||
from ..metadata_collector import get_metadata # type: ignore[import-not-found]
|
||||
from ..metadata_collector.metadata_processor import ( # type: ignore[import-not-found]
|
||||
from ..metadata_collector import get_metadata # pyright: ignore[reportMissingImports]
|
||||
from ..metadata_collector.metadata_processor import ( # pyright: ignore[reportMissingImports]
|
||||
MetadataProcessor,
|
||||
)
|
||||
from ..metadata_collector.metadata_registry import ( # type: ignore[import-not-found]
|
||||
from ..metadata_collector.metadata_registry import ( # pyright: ignore[reportMissingImports]
|
||||
MetadataRegistry,
|
||||
)
|
||||
else: # pragma: no cover - optional dependency path
|
||||
get_metadata = None # type: ignore[assignment]
|
||||
MetadataProcessor = None # type: ignore[assignment]
|
||||
MetadataRegistry = None # type: ignore[assignment]
|
||||
get_metadata = None # pyright: ignore[reportAssignmentType]
|
||||
MetadataProcessor = None # pyright: ignore[reportAssignmentType]
|
||||
MetadataRegistry = None # pyright: ignore[reportAssignmentType]
|
||||
|
||||
analysis_service = RecipeAnalysisService(
|
||||
exif_utils=ExifUtils,
|
||||
@@ -196,6 +201,18 @@ class BaseRecipeRoutes:
|
||||
sharing_service=sharing_service,
|
||||
)
|
||||
|
||||
# Lazy import: standalone mode replaces the ``server`` module with a
|
||||
# mock, so resolve PromptServer at handler-set build time instead of
|
||||
# module import time. The handler's standalone check guards UX.
|
||||
from server import PromptServer # pyright: ignore[reportMissingImports]
|
||||
|
||||
workflow = RecipeWorkflowHandler(
|
||||
ensure_dependencies_ready=self.ensure_dependencies_ready,
|
||||
recipe_scanner_getter=recipe_scanner_getter,
|
||||
prompt_server=PromptServer,
|
||||
logger=logger,
|
||||
)
|
||||
|
||||
from ..services.websocket_manager import ws_manager
|
||||
|
||||
batch_import_service = BatchImportService(
|
||||
@@ -220,4 +237,5 @@ class BaseRecipeRoutes:
|
||||
analysis=analysis,
|
||||
sharing=sharing,
|
||||
batch_import=batch_import,
|
||||
workflow=workflow,
|
||||
)
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import logging
|
||||
from typing import Dict, List, Set
|
||||
import os
|
||||
from typing import Any, Dict, List, Set
|
||||
from aiohttp import web
|
||||
|
||||
from .base_model_routes import BaseModelRoutes
|
||||
@@ -7,6 +8,7 @@ from .model_route_registrar import ModelRouteRegistrar
|
||||
from ..services.checkpoint_service import CheckpointService
|
||||
from ..services.service_registry import ServiceRegistry
|
||||
from ..config import config
|
||||
from ..utils.utils import _format_model_name_for_comfyui
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -28,13 +30,13 @@ class CheckpointRoutes(BaseModelRoutes):
|
||||
# Attach service dependencies
|
||||
self.attach_service(self.service)
|
||||
|
||||
def setup_routes(self, app: web.Application):
|
||||
def setup_routes(self, app: web.Application, prefix: str = "checkpoints"):
|
||||
"""Setup Checkpoint routes"""
|
||||
# Schedule service initialization on app startup
|
||||
app.on_startup.append(lambda _: self.initialize_services())
|
||||
|
||||
# Setup common routes with 'checkpoints' prefix (includes page route)
|
||||
super().setup_routes(app, 'checkpoints')
|
||||
super().setup_routes(app, prefix)
|
||||
|
||||
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
|
||||
"""Setup Checkpoint-specific routes"""
|
||||
@@ -45,6 +47,45 @@ 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}/unet_roots', prefix, self.get_unet_roots)
|
||||
|
||||
# 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)
|
||||
|
||||
async def get_loader_pool(self, request: web.Request) -> web.Response:
|
||||
"""Return ComfyUI-formatted model names with their base_model.
|
||||
|
||||
Backing data for the Checkpoint/Unet Loader nodes'
|
||||
control_after_generate feature: the front-end filters the
|
||||
ckpt_name/unet_name combo options by base_model using this pool, so
|
||||
randomize mode picks within the narrowed set.
|
||||
"""
|
||||
try:
|
||||
sub_type = request.query.get("sub_type", "checkpoint")
|
||||
if sub_type not in ("checkpoint", "diffusion_model"):
|
||||
return web.json_response({"error": "invalid sub_type"}, status=400)
|
||||
scanner = await ServiceRegistry.get_checkpoint_scanner()
|
||||
cache = await scanner.get_cached_data()
|
||||
model_roots = scanner.get_model_roots()
|
||||
items: List[Dict[str, str]] = []
|
||||
for item in cache.raw_data:
|
||||
if item.get("sub_type") != sub_type:
|
||||
continue
|
||||
file_path = item.get("file_path", "")
|
||||
if not file_path or not os.path.exists(file_path):
|
||||
continue
|
||||
formatted_name = _format_model_name_for_comfyui(file_path, model_roots)
|
||||
if formatted_name:
|
||||
items.append(
|
||||
{
|
||||
"name": formatted_name,
|
||||
"base_model": item.get("base_model", "") or "",
|
||||
}
|
||||
)
|
||||
items.sort(key=lambda x: x["name"])
|
||||
return web.json_response({"items": items})
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting loader pool: {e}", exc_info=True)
|
||||
return web.json_response({"error": str(e)}, status=500)
|
||||
|
||||
def _validate_civitai_model_type(self, model_type: str) -> bool:
|
||||
"""Validate CivitAI model type for Checkpoint"""
|
||||
return model_type.lower() == 'checkpoint'
|
||||
@@ -53,9 +94,9 @@ class CheckpointRoutes(BaseModelRoutes):
|
||||
"""Get expected model types string for error messages"""
|
||||
return "Checkpoint"
|
||||
|
||||
def _parse_specific_params(self, request: web.Request) -> Dict:
|
||||
def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]:
|
||||
"""Parse Checkpoint-specific parameters"""
|
||||
params: Dict = {}
|
||||
params: Dict[str, Any] = {}
|
||||
|
||||
if 'checkpoint_hash' in request.query:
|
||||
params['hash_filters'] = {'single_hash': request.query['checkpoint_hash'].lower()}
|
||||
@@ -70,7 +111,7 @@ class CheckpointRoutes(BaseModelRoutes):
|
||||
"""Get detailed information for a specific checkpoint by name"""
|
||||
try:
|
||||
name = request.match_info.get('name', '')
|
||||
checkpoint_info = await self.service.get_model_info_by_name(name)
|
||||
checkpoint_info = await self.service.get_model_info_by_name(name) # pyright: ignore[reportAttributeAccessIssue]
|
||||
|
||||
if checkpoint_info:
|
||||
return web.json_response(checkpoint_info)
|
||||
@@ -89,7 +130,7 @@ class CheckpointRoutes(BaseModelRoutes):
|
||||
roots.extend(config.checkpoints_roots or [])
|
||||
roots.extend(config.extra_checkpoints_roots or [])
|
||||
# Remove duplicates while preserving order
|
||||
seen: set = set()
|
||||
seen: set[str] = set()
|
||||
unique_roots: List[str] = []
|
||||
for root in roots:
|
||||
if root and root not in seen:
|
||||
@@ -114,7 +155,7 @@ class CheckpointRoutes(BaseModelRoutes):
|
||||
roots.extend(config.unet_roots or [])
|
||||
roots.extend(config.extra_unet_roots or [])
|
||||
# Remove duplicates while preserving order
|
||||
seen: set = set()
|
||||
seen: set[str] = set()
|
||||
unique_roots: List[str] = []
|
||||
for root in roots:
|
||||
if root and root not in seen:
|
||||
|
||||
@@ -26,13 +26,13 @@ class EmbeddingRoutes(BaseModelRoutes):
|
||||
# Attach service dependencies
|
||||
self.attach_service(self.service)
|
||||
|
||||
def setup_routes(self, app: web.Application):
|
||||
def setup_routes(self, app: web.Application, prefix: str = "embeddings"):
|
||||
"""Setup Embedding routes"""
|
||||
# Schedule service initialization on app startup
|
||||
app.on_startup.append(lambda _: self.initialize_services())
|
||||
|
||||
# Setup common routes with 'embeddings' prefix (includes page route)
|
||||
super().setup_routes(app, 'embeddings')
|
||||
super().setup_routes(app, prefix)
|
||||
|
||||
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
|
||||
"""Setup Embedding-specific routes"""
|
||||
@@ -51,7 +51,7 @@ class EmbeddingRoutes(BaseModelRoutes):
|
||||
"""Get detailed information for a specific embedding by name"""
|
||||
try:
|
||||
name = request.match_info.get('name', '')
|
||||
embedding_info = await self.service.get_model_info_by_name(name)
|
||||
embedding_info = await self.service.get_model_info_by_name(name) # pyright: ignore[reportAttributeAccessIssue]
|
||||
|
||||
if embedding_info:
|
||||
return web.json_response(embedding_info)
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Callable, Mapping
|
||||
from typing import Any, Awaitable, Callable, Mapping
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
@@ -35,7 +35,7 @@ class ExampleImagesRoutes:
|
||||
*,
|
||||
ws_manager,
|
||||
download_manager: DownloadManager | None = None,
|
||||
processor=ExampleImagesProcessor,
|
||||
processor: Any = ExampleImagesProcessor,
|
||||
file_manager=ExampleImagesFileManager,
|
||||
cleanup_service: ExampleImagesCleanupService | None = None,
|
||||
) -> None:
|
||||
@@ -46,7 +46,9 @@ class ExampleImagesRoutes:
|
||||
self._file_manager = file_manager
|
||||
self._cleanup_service = cleanup_service or ExampleImagesCleanupService()
|
||||
self._handler_set: ExampleImagesHandlerSet | None = None
|
||||
self._handler_mapping: Mapping[str, Callable[[web.Request], web.StreamResponse]] | None = None
|
||||
self._handler_mapping: Mapping[
|
||||
str, Callable[[web.Request], Awaitable[web.StreamResponse]]
|
||||
] | None = None
|
||||
|
||||
@classmethod
|
||||
def setup_routes(cls, app: web.Application, *, ws_manager) -> None:
|
||||
@@ -61,7 +63,9 @@ class ExampleImagesRoutes:
|
||||
registrar = ExampleImagesRouteRegistrar(app)
|
||||
registrar.register_routes(self.to_route_mapping())
|
||||
|
||||
def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
|
||||
def to_route_mapping(
|
||||
self,
|
||||
) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
|
||||
"""Return the registrar-compatible mapping of handler names to callables."""
|
||||
|
||||
if self._handler_mapping is None:
|
||||
|
||||
@@ -3,7 +3,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Mapping
|
||||
from typing import Awaitable, Callable, Mapping
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
@@ -170,7 +170,7 @@ class ExampleImagesHandlerSet:
|
||||
management: ExampleImagesManagementHandler
|
||||
files: ExampleImagesFileHandler
|
||||
|
||||
def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
|
||||
def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
|
||||
"""Flatten handler methods into the registrar mapping."""
|
||||
|
||||
return {
|
||||
|
||||
@@ -56,6 +56,7 @@ from ...utils.constants import (
|
||||
)
|
||||
from .hf_handlers import HfHandler
|
||||
from .agent_handlers import AgentHandler
|
||||
from .model_handlers import ModelCivitaiHandler
|
||||
from ...utils.civitai_utils import rewrite_preview_url
|
||||
from ...utils.example_images_paths import (
|
||||
find_non_compliant_items_in_example_images_root,
|
||||
@@ -276,7 +277,7 @@ def _collect_comfyui_session_logs(
|
||||
) -> dict[str, Any]:
|
||||
if log_entries is None:
|
||||
try:
|
||||
import app.logger as comfy_logger
|
||||
import app.logger as comfy_logger # pyright: ignore[reportMissingImports]
|
||||
|
||||
log_entries = list(comfy_logger.get_logs() or [])
|
||||
except Exception as exc: # pragma: no cover - environment dependent
|
||||
@@ -422,10 +423,10 @@ class PromptServerProtocol(Protocol):
|
||||
"""Subset of PromptServer used by the handlers."""
|
||||
|
||||
instance: "PromptServerProtocol"
|
||||
sockets: dict # maps clientId (sid) → WebSocketResponse
|
||||
sockets: dict[str, Any] # maps clientId (sid) → WebSocketResponse
|
||||
|
||||
def send_sync(
|
||||
self, event: str, payload: dict | None = None, sid: str | None = None
|
||||
self, event: str, payload: dict[str, Any] | None = None, sid: str | None = None
|
||||
) -> None: # pragma: no cover - protocol
|
||||
...
|
||||
|
||||
@@ -443,7 +444,12 @@ class UsageStatsFactory(Protocol):
|
||||
class MetadataProviderProtocol(Protocol):
|
||||
async def get_model_versions(
|
||||
self, model_id: int
|
||||
) -> dict | None: # pragma: no cover - protocol
|
||||
) -> dict[str, Any] | None: # pragma: no cover - protocol
|
||||
...
|
||||
|
||||
async def get_user_models(
|
||||
self, username: str, cursor: str | None = None
|
||||
) -> Any: # pragma: no cover - protocol
|
||||
...
|
||||
|
||||
|
||||
@@ -466,16 +472,16 @@ class MetadataArchiveManagerProtocol(Protocol):
|
||||
class BackupServiceProtocol(Protocol):
|
||||
async def create_snapshot(
|
||||
self, *, snapshot_type: str = "manual", persist: bool = False
|
||||
) -> dict: # pragma: no cover - protocol
|
||||
) -> dict[str, Any]: # pragma: no cover - protocol
|
||||
...
|
||||
|
||||
async def restore_snapshot(self, archive_path: str) -> dict: # pragma: no cover - protocol
|
||||
async def restore_snapshot(self, archive_path: str) -> dict[str, Any]: # pragma: no cover - protocol
|
||||
...
|
||||
|
||||
def get_status(self) -> dict: # pragma: no cover - protocol
|
||||
def get_status(self) -> dict[str, Any]: # pragma: no cover - protocol
|
||||
...
|
||||
|
||||
def get_available_snapshots(self) -> list[dict]: # pragma: no cover - protocol
|
||||
def get_available_snapshots(self) -> list[dict[str, Any]]: # pragma: no cover - protocol
|
||||
...
|
||||
|
||||
|
||||
@@ -491,7 +497,7 @@ class NodeRegistry:
|
||||
def __init__(self) -> None:
|
||||
self._lock = asyncio.Lock()
|
||||
# sid → {unique_id → node_info}
|
||||
self._tab_nodes: Dict[str, Dict[str, dict]] = {}
|
||||
self._tab_nodes: Dict[str, Dict[str, dict[str, Any]]] = {}
|
||||
self._ready = asyncio.Event()
|
||||
self._waiting_clients: set[str] = set()
|
||||
|
||||
@@ -504,7 +510,7 @@ class NodeRegistry:
|
||||
# Helpers to build one node dict (extracted so it's reused for each tab)
|
||||
# ------------------------------------------------------------------
|
||||
@staticmethod
|
||||
def _build_node_dict(node: dict) -> dict:
|
||||
def _build_node_dict(node: dict[str, Any]) -> dict[str, Any]:
|
||||
node_id = node["node_id"]
|
||||
graph_id = str(node["graph_id"])
|
||||
unique_id = f"{graph_id}:{node_id}"
|
||||
@@ -513,11 +519,11 @@ class NodeRegistry:
|
||||
bgcolor = node.get("bgcolor") or DEFAULT_NODE_COLOR
|
||||
|
||||
raw_capabilities = node.get("capabilities")
|
||||
capabilities: dict = {}
|
||||
capabilities: dict[str, Any] = {}
|
||||
if isinstance(raw_capabilities, dict):
|
||||
capabilities = dict(raw_capabilities)
|
||||
|
||||
raw_widget_names: list | None = node.get("widget_names")
|
||||
raw_widget_names: list[Any] | None = node.get("widget_names")
|
||||
if not isinstance(raw_widget_names, list):
|
||||
capability_widget_names = capabilities.get("widget_names")
|
||||
raw_widget_names = (
|
||||
@@ -565,9 +571,9 @@ class NodeRegistry:
|
||||
# ------------------------------------------------------------------
|
||||
# Public API
|
||||
# ------------------------------------------------------------------
|
||||
async def register_nodes(self, sid: str, nodes: list[dict]) -> None:
|
||||
async def register_nodes(self, sid: str, nodes: list[dict[str, Any]]) -> None:
|
||||
"""Register/replace the node list for a single ComfyUI tab (identified by *sid*)."""
|
||||
tab_nodes: dict[str, dict] = {}
|
||||
tab_nodes: dict[str, dict[str, Any]] = {}
|
||||
for node in nodes:
|
||||
nd = self._build_node_dict(node)
|
||||
tab_nodes[nd["unique_id"]] = nd
|
||||
@@ -602,7 +608,7 @@ class NodeRegistry:
|
||||
except asyncio.TimeoutError:
|
||||
return False
|
||||
|
||||
async def get_merged_registry(self, active_sids: set[str] | None = None) -> dict:
|
||||
async def get_merged_registry(self, active_sids: set[str] | None = None) -> dict[str, Any]:
|
||||
"""Return the union of all known tab nodes, pruning any tab that is no
|
||||
longer connected."""
|
||||
async with self._lock:
|
||||
@@ -619,8 +625,8 @@ class NodeRegistry:
|
||||
len(stale_sids), stale_sids,
|
||||
)
|
||||
|
||||
merged: dict[str, dict] = {}
|
||||
tab_info: dict[str, dict] = {}
|
||||
merged: dict[str, dict[str, Any]] = {}
|
||||
tab_info: dict[str, dict[str, Any]] = {}
|
||||
for sid, nodes in self._tab_nodes.items():
|
||||
tab_info[sid] = {
|
||||
"node_count": len(nodes),
|
||||
@@ -643,9 +649,60 @@ class NodeRegistry:
|
||||
|
||||
|
||||
class HealthCheckHandler:
|
||||
def __init__(
|
||||
self,
|
||||
scanner_getters: Mapping[str, Callable[[], Awaitable[Any]]] | None = None,
|
||||
) -> None:
|
||||
self._scanner_getters = scanner_getters or {
|
||||
"lora": ServiceRegistry.get_lora_scanner,
|
||||
"checkpoint": ServiceRegistry.get_checkpoint_scanner,
|
||||
"embedding": ServiceRegistry.get_embedding_scanner,
|
||||
"recipe": ServiceRegistry.get_recipe_scanner,
|
||||
}
|
||||
|
||||
async def health_check(self, request: web.Request) -> web.Response:
|
||||
return web.json_response({"status": "ok"})
|
||||
|
||||
async def get_init_status(self, request: web.Request) -> web.Response:
|
||||
"""Report aggregate scanner initialization status.
|
||||
|
||||
Used by the initialization page's polling fallback when the
|
||||
/ws/init-progress WebSocket is unavailable. Omits pageType so every
|
||||
page accepts the update and only reloads once all scanners are done.
|
||||
"""
|
||||
pending: list[str] = []
|
||||
for name, getter in self._scanner_getters.items():
|
||||
try:
|
||||
scanner = await getter()
|
||||
except Exception:
|
||||
pending.append(name)
|
||||
continue
|
||||
cache_ready = getattr(scanner, "_cache", None) is not None
|
||||
is_initializing = getattr(scanner, "is_initializing", None)
|
||||
busy = (
|
||||
is_initializing()
|
||||
if callable(is_initializing)
|
||||
else bool(getattr(scanner, "_is_initializing", False))
|
||||
)
|
||||
if busy or not cache_ready:
|
||||
pending.append(name)
|
||||
|
||||
if pending:
|
||||
return web.json_response(
|
||||
{
|
||||
"status": "initializing",
|
||||
"stage": "processing",
|
||||
"details": "Initializing: " + ", ".join(pending),
|
||||
}
|
||||
)
|
||||
return web.json_response(
|
||||
{
|
||||
"status": "complete",
|
||||
"progress": 100,
|
||||
"details": "Initialization complete",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class SupportersHandler:
|
||||
"""Handler for supporters data."""
|
||||
@@ -653,7 +710,7 @@ class SupportersHandler:
|
||||
def __init__(self, logger: logging.Logger | None = None) -> None:
|
||||
self._logger = logger or logging.getLogger(__name__)
|
||||
|
||||
def _load_supporters(self) -> dict:
|
||||
def _load_supporters(self) -> dict[str, Any]:
|
||||
"""Load supporters data from JSON file."""
|
||||
try:
|
||||
current_file = os.path.abspath(__file__)
|
||||
@@ -1229,10 +1286,8 @@ class DoctorHandler:
|
||||
settings_snapshot = _sanitize_sensitive_data(
|
||||
getattr(self._settings, "settings", {}) or {}
|
||||
)
|
||||
startup_messages_getter = getattr(self._settings, "get_startup_messages", None)
|
||||
startup_messages = (
|
||||
list(startup_messages_getter()) if callable(startup_messages_getter) else []
|
||||
)
|
||||
startup_messages_getter: Any = getattr(self._settings, "get_startup_messages", None)
|
||||
startup_messages = list(startup_messages_getter()) if startup_messages_getter else []
|
||||
|
||||
environment = {
|
||||
"app_version": app_version,
|
||||
@@ -1439,7 +1494,7 @@ class SettingsHandler:
|
||||
*,
|
||||
settings_service=None,
|
||||
metadata_provider_updater: Callable[
|
||||
[], Awaitable[None]
|
||||
[], Awaitable[Any]
|
||||
] = update_metadata_providers,
|
||||
downloader_factory: Callable[
|
||||
[], Awaitable[DownloaderProtocol]
|
||||
@@ -1484,8 +1539,8 @@ class SettingsHandler:
|
||||
settings_file = getattr(self._settings, "settings_file", None)
|
||||
if settings_file:
|
||||
response_data["settings_file"] = settings_file
|
||||
messages_getter = getattr(self._settings, "get_startup_messages", None)
|
||||
messages = list(messages_getter()) if callable(messages_getter) else []
|
||||
messages_getter: Any = getattr(self._settings, "get_startup_messages", None)
|
||||
messages = list(messages_getter()) if messages_getter else []
|
||||
return web.json_response(
|
||||
{
|
||||
"success": True,
|
||||
@@ -1562,6 +1617,11 @@ class SettingsHandler:
|
||||
{"success": False, "error": validation_error}
|
||||
)
|
||||
|
||||
if key == "update_channel" and value not in ("release", "nightly"):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "update_channel must be 'release' or 'nightly'"}
|
||||
)
|
||||
|
||||
if value == "__DELETE__" and key in (
|
||||
"proxy_username",
|
||||
"proxy_password",
|
||||
@@ -1570,7 +1630,11 @@ class SettingsHandler:
|
||||
else:
|
||||
self._settings.set(key, value)
|
||||
|
||||
if key == "enable_metadata_archive_db":
|
||||
if key in (
|
||||
"enable_metadata_archive_db",
|
||||
"enable_civarchive_api",
|
||||
"metadata_provider_order",
|
||||
):
|
||||
await self._metadata_provider_updater()
|
||||
|
||||
if key in self._PROXY_KEYS:
|
||||
@@ -1784,6 +1848,124 @@ class LoraCodeHandler:
|
||||
logger.error("Failed to update lora code: %s", exc, exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||
|
||||
async def get_update_lora_code(self, request: web.Request) -> web.Response:
|
||||
"""GET version of update_lora_code — reads parameters from query string.
|
||||
|
||||
Query params:
|
||||
lora_code (required) — the LoRA syntax to send
|
||||
mode (optional) — "append" (default) or "replace"
|
||||
node_id (repeatable) — target node id(s), e.g. node_id=3&node_id=5
|
||||
node_ids (optional) — JSON-encoded array for complex references with graph_id:
|
||||
[{"node_id":3,"graph_id":"g1"}, ...]
|
||||
"""
|
||||
try:
|
||||
node_ids_raw = request.query.get("node_ids")
|
||||
node_id_list = request.query.getall("node_id", [])
|
||||
lora_code = request.query.get("lora_code", "")
|
||||
mode = request.query.get("mode", "append")
|
||||
|
||||
if not lora_code:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Missing lora_code parameter"},
|
||||
status=400,
|
||||
)
|
||||
|
||||
node_ids = None
|
||||
if node_ids_raw:
|
||||
try:
|
||||
node_ids = json.loads(node_ids_raw)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "node_ids must be a valid JSON array"},
|
||||
status=400,
|
||||
)
|
||||
if not isinstance(node_ids, list) or not node_ids:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "node_ids must be a non-empty JSON array"},
|
||||
status=400,
|
||||
)
|
||||
elif node_id_list:
|
||||
node_ids = node_id_list
|
||||
|
||||
results = []
|
||||
if node_ids is None:
|
||||
try:
|
||||
self._prompt_server.instance.send_sync(
|
||||
"lora_code_update",
|
||||
{"id": -1, "lora_code": lora_code, "mode": mode},
|
||||
)
|
||||
results.append({"node_id": "broadcast", "success": True})
|
||||
except Exception as exc: # pragma: no cover - defensive logging
|
||||
logger.error("Error broadcasting lora code: %s", exc)
|
||||
results.append(
|
||||
{"node_id": "broadcast", "success": False, "error": str(exc)}
|
||||
)
|
||||
else:
|
||||
for entry in node_ids:
|
||||
node_identifier = entry
|
||||
graph_identifier = None
|
||||
if isinstance(entry, dict):
|
||||
node_identifier = entry.get("node_id")
|
||||
graph_identifier = entry.get("graph_id")
|
||||
|
||||
if node_identifier is None:
|
||||
results.append(
|
||||
{
|
||||
"node_id": node_identifier,
|
||||
"graph_id": graph_identifier,
|
||||
"success": False,
|
||||
"error": "Missing node_id parameter",
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
try:
|
||||
parsed_node_id = int(node_identifier)
|
||||
except (TypeError, ValueError):
|
||||
parsed_node_id = node_identifier
|
||||
|
||||
payload = {
|
||||
"id": parsed_node_id,
|
||||
"lora_code": lora_code,
|
||||
"mode": mode,
|
||||
}
|
||||
|
||||
if graph_identifier is not None:
|
||||
payload["graph_id"] = str(graph_identifier)
|
||||
|
||||
try:
|
||||
self._prompt_server.instance.send_sync(
|
||||
"lora_code_update",
|
||||
payload,
|
||||
)
|
||||
results.append(
|
||||
{
|
||||
"node_id": parsed_node_id,
|
||||
"graph_id": payload.get("graph_id"),
|
||||
"success": True,
|
||||
}
|
||||
)
|
||||
except Exception as exc: # pragma: no cover - defensive logging
|
||||
logger.error(
|
||||
"Error sending lora code to node %s (graph %s): %s",
|
||||
parsed_node_id,
|
||||
graph_identifier,
|
||||
exc,
|
||||
)
|
||||
results.append(
|
||||
{
|
||||
"node_id": parsed_node_id,
|
||||
"graph_id": payload.get("graph_id"),
|
||||
"success": False,
|
||||
"error": str(exc),
|
||||
}
|
||||
)
|
||||
|
||||
return web.json_response({"success": True, "results": results})
|
||||
except Exception as exc: # pragma: no cover - defensive logging
|
||||
logger.error("Failed to update lora code (GET): %s", exc, exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||
|
||||
|
||||
class TrainedWordsHandler:
|
||||
async def get_trained_words(self, request: web.Request) -> web.Response:
|
||||
@@ -1878,11 +2060,11 @@ async def _noop_backup_service() -> None:
|
||||
|
||||
@dataclass
|
||||
class ServiceRegistryAdapter:
|
||||
get_lora_scanner: Callable[[], Awaitable]
|
||||
get_checkpoint_scanner: Callable[[], Awaitable]
|
||||
get_embedding_scanner: Callable[[], Awaitable]
|
||||
get_downloaded_version_history_service: Callable[[], Awaitable]
|
||||
get_backup_service: Callable[[], Awaitable] = _noop_backup_service
|
||||
get_lora_scanner: Callable[[], Awaitable[Any]]
|
||||
get_checkpoint_scanner: Callable[[], Awaitable[Any]]
|
||||
get_embedding_scanner: Callable[[], Awaitable[Any]]
|
||||
get_downloaded_version_history_service: Callable[[], Awaitable[Any]]
|
||||
get_backup_service: Callable[[], Awaitable[Any]] = _noop_backup_service
|
||||
|
||||
|
||||
class ModelLibraryHandler:
|
||||
@@ -1923,14 +2105,71 @@ class ModelLibraryHandler:
|
||||
return await self._service_registry.get_downloaded_version_history_service()
|
||||
|
||||
@staticmethod
|
||||
def _with_downloaded_flag(versions: list[dict]) -> list[dict]:
|
||||
enriched: list[dict] = []
|
||||
def _with_downloaded_flag(versions: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
enriched: list[dict[str, Any]] = []
|
||||
for version in versions:
|
||||
entry = dict(version)
|
||||
entry.setdefault("hasBeenDownloaded", True)
|
||||
enriched.append(entry)
|
||||
return enriched
|
||||
|
||||
@staticmethod
|
||||
async def _get_downloaded_files(
|
||||
scanner: Any, model_version_id: int
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Return per-file downloaded state for a version in the library.
|
||||
|
||||
This handler has no CivitAI version payload, so the remote file list
|
||||
is taken from the local entries' cached ``civitai`` metadata (the
|
||||
full version payload persisted at download time, see
|
||||
``BaseModelMetadata.from_civitai_info``) and matched with the same
|
||||
D2 rule used by ``get_civitai_versions`` (#1058). Local entries that
|
||||
cannot be matched to a known remote file (e.g. missing metadata or
|
||||
renamed files) are still reported with ``fileId`` set to None.
|
||||
Returns ``[{fileId, fileName, filePath}]``.
|
||||
"""
|
||||
try:
|
||||
cache = await scanner.get_cached_data()
|
||||
except Exception: # pragma: no cover - defensive fallback
|
||||
logger.debug(
|
||||
"Failed to read cache for downloaded files of version %s",
|
||||
model_version_id,
|
||||
exc_info=True,
|
||||
)
|
||||
return []
|
||||
|
||||
files_getter = getattr(cache, "get_files_by_version_id", None)
|
||||
local_entries = files_getter(model_version_id) if files_getter else []
|
||||
if not local_entries:
|
||||
return []
|
||||
|
||||
version_payload: Mapping[str, Any] = {}
|
||||
for entry in local_entries:
|
||||
civitai = entry.get("civitai") if isinstance(entry, Mapping) else None
|
||||
if isinstance(civitai, Mapping) and isinstance(civitai.get("files"), list):
|
||||
version_payload = civitai
|
||||
break
|
||||
|
||||
downloaded = ModelCivitaiHandler._match_downloaded_files(
|
||||
version_payload, local_entries
|
||||
)
|
||||
|
||||
# Surface local files that D2 could not map to a known remote file
|
||||
matched_paths = {item.get("filePath") for item in downloaded}
|
||||
for entry in local_entries:
|
||||
if not isinstance(entry, Mapping):
|
||||
continue
|
||||
if entry.get("file_path") in matched_paths:
|
||||
continue
|
||||
downloaded.append(
|
||||
{
|
||||
"fileId": None,
|
||||
"fileName": entry.get("file_name"),
|
||||
"filePath": entry.get("file_path"),
|
||||
}
|
||||
)
|
||||
return downloaded
|
||||
|
||||
async def check_model_exists(self, request: web.Request) -> web.Response:
|
||||
try:
|
||||
model_id_str = request.query.get("modelId")
|
||||
@@ -1966,9 +2205,11 @@ class ModelLibraryHandler:
|
||||
|
||||
exists = False
|
||||
model_type = None
|
||||
matched_scanner = None
|
||||
if await lora_scanner.check_model_version_exists(model_version_id):
|
||||
exists = True
|
||||
model_type = "lora"
|
||||
matched_scanner = lora_scanner
|
||||
elif (
|
||||
checkpoint_scanner
|
||||
and await checkpoint_scanner.check_model_version_exists(
|
||||
@@ -1977,6 +2218,7 @@ class ModelLibraryHandler:
|
||||
):
|
||||
exists = True
|
||||
model_type = "checkpoint"
|
||||
matched_scanner = checkpoint_scanner
|
||||
elif (
|
||||
embedding_scanner
|
||||
and await embedding_scanner.check_model_version_exists(
|
||||
@@ -1985,6 +2227,7 @@ class ModelLibraryHandler:
|
||||
):
|
||||
exists = True
|
||||
model_type = "embedding"
|
||||
matched_scanner = embedding_scanner
|
||||
|
||||
if exists:
|
||||
return web.json_response(
|
||||
@@ -1993,6 +2236,9 @@ class ModelLibraryHandler:
|
||||
"exists": True,
|
||||
"modelType": model_type,
|
||||
"hasBeenDownloaded": False,
|
||||
"downloadedFiles": await self._get_downloaded_files(
|
||||
matched_scanner, model_version_id
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
@@ -2014,6 +2260,7 @@ class ModelLibraryHandler:
|
||||
"exists": False,
|
||||
"modelType": history_type,
|
||||
"hasBeenDownloaded": has_been_downloaded,
|
||||
"downloadedFiles": [],
|
||||
}
|
||||
)
|
||||
|
||||
@@ -2117,7 +2364,7 @@ class ModelLibraryHandler:
|
||||
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
||||
embedding_scanner = await self._service_registry.get_embedding_scanner()
|
||||
|
||||
results: list[dict] = []
|
||||
results: list[dict[str, Any]] = []
|
||||
for model_id in model_ids:
|
||||
lora_versions = await lora_scanner.get_model_versions_by_id(model_id)
|
||||
if lora_versions:
|
||||
@@ -2226,7 +2473,7 @@ class ModelLibraryHandler:
|
||||
)
|
||||
|
||||
try:
|
||||
model_version_id = int(data.get("modelVersionId"))
|
||||
model_version_id = int(data.get("modelVersionId")) # pyright: ignore[reportArgumentType]
|
||||
except (TypeError, ValueError):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Parameter modelVersionId must be an integer"},
|
||||
@@ -2298,8 +2545,8 @@ class ModelLibraryHandler:
|
||||
embedding_scanner = await self._service_registry.get_embedding_scanner()
|
||||
|
||||
found_type = None
|
||||
file_path = None
|
||||
found_cache = None
|
||||
entries: list = []
|
||||
|
||||
for model_type, scanner in (
|
||||
("lora", lora_scanner),
|
||||
@@ -2310,27 +2557,43 @@ class ModelLibraryHandler:
|
||||
if cache and model_version_id in cache.version_index:
|
||||
found_type = model_type
|
||||
found_cache = cache
|
||||
entry = cache.version_index[model_version_id]
|
||||
file_path = entry.get("file_path")
|
||||
# A version can have several local files (#1058); collect
|
||||
# them all so the delete below covers every file.
|
||||
files_getter = getattr(cache, "get_files_by_version_id", None)
|
||||
if files_getter is not None:
|
||||
entries = files_getter(model_version_id)
|
||||
else:
|
||||
entries = [cache.version_index[model_version_id]]
|
||||
break
|
||||
|
||||
if not file_path:
|
||||
file_paths = [
|
||||
entry.get("file_path")
|
||||
for entry in entries
|
||||
if isinstance(entry, dict) and entry.get("file_path")
|
||||
]
|
||||
|
||||
if not file_paths:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Model version not found in any scanner cache"},
|
||||
status=404,
|
||||
)
|
||||
|
||||
for file_path in file_paths:
|
||||
target_dir = os.path.dirname(file_path)
|
||||
base_name = os.path.basename(file_path)
|
||||
file_name, extension = os.path.splitext(base_name)
|
||||
await delete_model_artifacts(target_dir, file_name, main_extension=extension)
|
||||
|
||||
if found_cache:
|
||||
removed_paths = set(file_paths)
|
||||
found_cache.raw_data = [
|
||||
item
|
||||
for item in found_cache.raw_data
|
||||
if item.get("file_path") != file_path
|
||||
if item.get("file_path") not in removed_paths
|
||||
]
|
||||
rebuild = getattr(found_cache, "rebuild_version_index", None)
|
||||
if rebuild is not None:
|
||||
rebuild()
|
||||
await found_cache.resort()
|
||||
|
||||
scanner_map = {
|
||||
@@ -2338,10 +2601,11 @@ class ModelLibraryHandler:
|
||||
"checkpoint": checkpoint_scanner,
|
||||
"embedding": embedding_scanner,
|
||||
}
|
||||
scanner = scanner_map.get(found_type)
|
||||
scanner = scanner_map.get(found_type or "")
|
||||
if scanner:
|
||||
persist = getattr(scanner, "_persist_current_cache", None)
|
||||
if callable(persist):
|
||||
scanner.bump_cache_version()
|
||||
persist: Any = getattr(scanner, "_persist_current_cache", None)
|
||||
if persist:
|
||||
await persist()
|
||||
|
||||
history_service = await self._get_download_history_service()
|
||||
@@ -2352,6 +2616,7 @@ class ModelLibraryHandler:
|
||||
"success": True,
|
||||
"modelType": found_type,
|
||||
"modelVersionId": model_version_id,
|
||||
"deletedFiles": len(file_paths),
|
||||
}
|
||||
)
|
||||
except Exception as exc:
|
||||
@@ -2463,6 +2728,8 @@ class ModelLibraryHandler:
|
||||
status=400,
|
||||
)
|
||||
|
||||
cursor = request.query.get("cursor")
|
||||
|
||||
metadata_provider = await self._metadata_provider_factory()
|
||||
if not metadata_provider:
|
||||
return web.json_response(
|
||||
@@ -2471,7 +2738,7 @@ class ModelLibraryHandler:
|
||||
)
|
||||
|
||||
try:
|
||||
models = await metadata_provider.get_user_models(username)
|
||||
result = await metadata_provider.get_user_models(username, cursor)
|
||||
except NotImplementedError:
|
||||
return web.json_response(
|
||||
{
|
||||
@@ -2481,14 +2748,35 @@ class ModelLibraryHandler:
|
||||
status=501,
|
||||
)
|
||||
|
||||
if models is None:
|
||||
if result is None:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Failed to fetch user models"},
|
||||
status=502,
|
||||
)
|
||||
|
||||
if isinstance(result, dict):
|
||||
models = result.get("items")
|
||||
next_cursor = result.get("nextCursor")
|
||||
else:
|
||||
# Defensive: tolerate providers that still return a raw list
|
||||
models = result
|
||||
next_cursor = None
|
||||
|
||||
if not isinstance(models, list):
|
||||
models = []
|
||||
if next_cursor is not None and not isinstance(next_cursor, str):
|
||||
next_cursor = str(next_cursor)
|
||||
|
||||
estimated_total = None
|
||||
if cursor is None:
|
||||
get_count = getattr(metadata_provider, "get_creator_model_count", None)
|
||||
if get_count is not None:
|
||||
try:
|
||||
estimated_total = await get_count(username)
|
||||
except Exception: # best-effort only
|
||||
estimated_total = None
|
||||
if not isinstance(estimated_total, int):
|
||||
estimated_total = None
|
||||
|
||||
lora_scanner = await self._service_registry.get_lora_scanner()
|
||||
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
||||
@@ -2499,15 +2787,16 @@ class ModelLibraryHandler:
|
||||
}
|
||||
lora_type_aliases = {model_type.lower() for model_type in VALID_LORA_TYPES}
|
||||
|
||||
type_scanner_map: Dict[str, object | None] = {
|
||||
type_scanner_map: Dict[str, Any] = {
|
||||
**{alias: lora_scanner for alias in lora_type_aliases},
|
||||
"checkpoint": checkpoint_scanner,
|
||||
"textualinversion": embedding_scanner,
|
||||
}
|
||||
|
||||
versions: list[dict] = []
|
||||
versions: list[dict[str, Any]] = []
|
||||
history_service = await self._get_download_history_service()
|
||||
model_ids: list[int] = []
|
||||
model_count = 0
|
||||
for model in models:
|
||||
try:
|
||||
model_ids.append(int(model.get("id")))
|
||||
@@ -2541,6 +2830,8 @@ class ModelLibraryHandler:
|
||||
if model_type not in normalized_allowed_types:
|
||||
continue
|
||||
|
||||
model_count += 1
|
||||
|
||||
scanner = type_scanner_map.get(model_type)
|
||||
if scanner is None:
|
||||
return web.json_response(
|
||||
@@ -2554,6 +2845,8 @@ class ModelLibraryHandler:
|
||||
tags_value = model.get("tags")
|
||||
tags = tags_value if isinstance(tags_value, list) else []
|
||||
model_id = model.get("id")
|
||||
if model_id is None:
|
||||
continue
|
||||
try:
|
||||
model_id_int = int(model_id)
|
||||
except (TypeError, ValueError):
|
||||
@@ -2569,6 +2862,8 @@ class ModelLibraryHandler:
|
||||
continue
|
||||
|
||||
version_id = version.get("id")
|
||||
if version_id is None:
|
||||
continue
|
||||
try:
|
||||
version_id_int = int(version_id)
|
||||
except (TypeError, ValueError):
|
||||
@@ -2606,7 +2901,15 @@ class ModelLibraryHandler:
|
||||
)
|
||||
|
||||
return web.json_response(
|
||||
{"success": True, "username": username, "versions": versions}
|
||||
{
|
||||
"success": True,
|
||||
"username": username,
|
||||
"versions": versions,
|
||||
"modelCount": model_count,
|
||||
"nextCursor": next_cursor,
|
||||
"hasMore": next_cursor is not None,
|
||||
"estimatedTotal": estimated_total,
|
||||
}
|
||||
)
|
||||
except Exception as exc: # pragma: no cover - defensive logging
|
||||
logger.error("Failed to get Civitai user models: %s", exc, exc_info=True)
|
||||
@@ -2622,7 +2925,7 @@ class MetadataArchiveHandler:
|
||||
] = get_metadata_archive_manager,
|
||||
settings_service=None,
|
||||
metadata_provider_updater: Callable[
|
||||
[], Awaitable[None]
|
||||
[], Awaitable[Any]
|
||||
] = update_metadata_providers,
|
||||
) -> None:
|
||||
self._metadata_archive_manager_factory = metadata_archive_manager_factory
|
||||
@@ -2769,7 +3072,7 @@ class BackupHandler:
|
||||
|
||||
if request.content_type.startswith("multipart/"):
|
||||
reader = await request.multipart()
|
||||
field = await reader.next()
|
||||
field: Any = await reader.next()
|
||||
uploaded = False
|
||||
while field is not None:
|
||||
if getattr(field, "filename", None):
|
||||
@@ -3353,7 +3656,7 @@ class NodeRegistryHandler:
|
||||
status=400,
|
||||
)
|
||||
|
||||
if not isinstance(value, str) or not value:
|
||||
if value is None or (isinstance(value, str) and not value):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Missing value parameter"}, status=400
|
||||
)
|
||||
@@ -3388,7 +3691,7 @@ class NodeRegistryHandler:
|
||||
except (TypeError, ValueError):
|
||||
parsed_node_id = node_identifier
|
||||
|
||||
payload: dict = {
|
||||
payload: dict[str, Any] = {
|
||||
"id": parsed_node_id,
|
||||
"value": value,
|
||||
"mode": mode,
|
||||
@@ -3431,6 +3734,130 @@ class NodeRegistryHandler:
|
||||
logger.error("Failed to update node widget: %s", exc, exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||
|
||||
async def get_update_node_widget(self, request: web.Request) -> web.Response:
|
||||
"""GET version of update_node_widget — reads parameters from query string.
|
||||
|
||||
Query params:
|
||||
widget_name (optional) — the widget name to update (required unless action is set)
|
||||
action (optional) — alternative action, e.g. "inject_text" (required unless widget_name is set)
|
||||
value (required) — the value to set
|
||||
mode (optional) — "replace" (default) or "append"
|
||||
node_id (repeatable) — target node id(s), e.g. node_id=3&node_id=5
|
||||
node_ids (optional) — JSON-encoded array for complex references:
|
||||
[{"node_id":3,"graph_id":"g1"}, ...]
|
||||
"""
|
||||
try:
|
||||
widget_name = request.query.get("widget_name")
|
||||
action = request.query.get("action")
|
||||
value = request.query.get("value")
|
||||
mode = request.query.get("mode", "replace")
|
||||
node_ids_raw = request.query.get("node_ids")
|
||||
node_id_list = request.query.getall("node_id", [])
|
||||
|
||||
if not action and (not isinstance(widget_name, str) or not widget_name):
|
||||
return web.json_response(
|
||||
{
|
||||
"success": False,
|
||||
"error": "Missing parameter: provide either 'action' or 'widget_name'",
|
||||
},
|
||||
status=400,
|
||||
)
|
||||
|
||||
if value is None or (isinstance(value, str) and not value):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Missing value parameter"}, status=400
|
||||
)
|
||||
|
||||
node_ids = None
|
||||
if node_ids_raw:
|
||||
try:
|
||||
node_ids = json.loads(node_ids_raw)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "node_ids must be a valid JSON array"},
|
||||
status=400,
|
||||
)
|
||||
if not isinstance(node_ids, list) or not node_ids:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "node_ids must be a non-empty JSON array"},
|
||||
status=400,
|
||||
)
|
||||
elif node_id_list:
|
||||
node_ids = node_id_list
|
||||
|
||||
if not isinstance(node_ids, list) or not node_ids:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "node_ids must be a non-empty list"},
|
||||
status=400,
|
||||
)
|
||||
|
||||
results = []
|
||||
for entry in node_ids:
|
||||
node_identifier = entry
|
||||
graph_identifier = None
|
||||
if isinstance(entry, dict):
|
||||
node_identifier = entry.get("node_id")
|
||||
graph_identifier = entry.get("graph_id")
|
||||
|
||||
if node_identifier is None:
|
||||
results.append(
|
||||
{
|
||||
"node_id": node_identifier,
|
||||
"graph_id": graph_identifier,
|
||||
"success": False,
|
||||
"error": "Missing node_id parameter",
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
try:
|
||||
parsed_node_id = int(node_identifier)
|
||||
except (TypeError, ValueError):
|
||||
parsed_node_id = node_identifier
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"id": parsed_node_id,
|
||||
"value": value,
|
||||
"mode": mode,
|
||||
}
|
||||
if action:
|
||||
payload["action"] = action
|
||||
if widget_name:
|
||||
payload["widget_name"] = widget_name
|
||||
|
||||
if graph_identifier is not None:
|
||||
payload["graph_id"] = str(graph_identifier)
|
||||
|
||||
try:
|
||||
self._prompt_server.instance.send_sync("lm_widget_update", payload)
|
||||
results.append(
|
||||
{
|
||||
"node_id": parsed_node_id,
|
||||
"graph_id": payload.get("graph_id"),
|
||||
"success": True,
|
||||
}
|
||||
)
|
||||
except Exception as exc: # pragma: no cover - defensive logging
|
||||
logger.error(
|
||||
"Error sending widget update to node %s (graph %s): %s",
|
||||
parsed_node_id,
|
||||
graph_identifier,
|
||||
exc,
|
||||
)
|
||||
results.append(
|
||||
{
|
||||
"node_id": parsed_node_id,
|
||||
"graph_id": payload.get("graph_id"),
|
||||
"success": False,
|
||||
"error": str(exc),
|
||||
}
|
||||
)
|
||||
|
||||
return web.json_response({"success": True, "results": results})
|
||||
except Exception as exc: # pragma: no cover - defensive logging
|
||||
logger.error("Failed to update node widget (GET): %s", exc, exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||
|
||||
|
||||
class MiscHandlerSet:
|
||||
"""Aggregate handlers into a lookup compatible with the registrar."""
|
||||
@@ -3455,8 +3882,8 @@ class MiscHandlerSet:
|
||||
doctor: DoctorHandler,
|
||||
example_workflows: ExampleWorkflowsHandler,
|
||||
base_model: BaseModelHandlerSet,
|
||||
hf_handler: HfHandler | None = None,
|
||||
agent_handler: AgentHandler | None = None,
|
||||
hf_handler: Any = None,
|
||||
agent_handler: Any = None,
|
||||
) -> None:
|
||||
self.health = health
|
||||
self.settings = settings
|
||||
@@ -3483,6 +3910,7 @@ class MiscHandlerSet:
|
||||
) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
|
||||
return {
|
||||
"health_check": self.health.health_check,
|
||||
"get_init_status": self.health.get_init_status,
|
||||
"get_settings": self.settings.get_settings,
|
||||
"update_settings": self.settings.update_settings,
|
||||
"get_doctor_diagnostics": self.doctor.get_doctor_diagnostics,
|
||||
@@ -3497,10 +3925,12 @@ class MiscHandlerSet:
|
||||
"update_usage_stats": self.usage_stats.update_usage_stats,
|
||||
"get_usage_stats": self.usage_stats.get_usage_stats,
|
||||
"update_lora_code": self.lora_code.update_lora_code,
|
||||
"get_update_lora_code": self.lora_code.get_update_lora_code,
|
||||
"get_trained_words": self.trained_words.get_trained_words,
|
||||
"get_model_example_files": self.model_examples.get_model_example_files,
|
||||
"register_nodes": self.node_registry.register_nodes,
|
||||
"update_node_widget": self.node_registry.update_node_widget,
|
||||
"get_update_node_widget": self.node_registry.get_update_node_widget,
|
||||
"get_registry": self.node_registry.get_registry,
|
||||
"check_model_exists": self.model_library.check_model_exists,
|
||||
"check_models_exist": self.model_library.check_models_exist,
|
||||
|
||||
@@ -15,6 +15,10 @@ from aiohttp import web
|
||||
import jinja2
|
||||
|
||||
from ...config import config
|
||||
from ...services.active_filters_store import (
|
||||
ActiveFiltersStore,
|
||||
active_filters_to_query_kwargs,
|
||||
)
|
||||
from ...services.download_coordinator import DownloadCoordinator
|
||||
from ...services.connectivity_guard import (
|
||||
OFFLINE_FRIENDLY_MESSAGE,
|
||||
@@ -51,6 +55,29 @@ LICENSE_FIELDS = (
|
||||
)
|
||||
|
||||
|
||||
_broadcast_models_changed_tasks: set = set()
|
||||
|
||||
|
||||
def _broadcast_models_changed() -> None:
|
||||
"""Notify connected clients that the local model library changed.
|
||||
|
||||
The ComfyUI graph page listens for this event to invalidate its cached
|
||||
model availability data (loras widget missing-model cues / error flags)
|
||||
without waiting for the cache TTL to expire.
|
||||
"""
|
||||
try:
|
||||
from ...services.websocket_manager import ws_manager
|
||||
|
||||
task = asyncio.create_task(ws_manager.broadcast({"type": "models_changed"}))
|
||||
# Keep a reference so the task is not garbage-collected mid-await.
|
||||
_broadcast_models_changed_tasks.add(task)
|
||||
task.add_done_callback(_broadcast_models_changed_tasks.discard)
|
||||
except Exception:
|
||||
logging.getLogger(__name__).debug(
|
||||
"Failed to broadcast models_changed", exc_info=True
|
||||
)
|
||||
|
||||
|
||||
class ModelPageView:
|
||||
"""Render the HTML view for model listings."""
|
||||
|
||||
@@ -71,7 +98,7 @@ class ModelPageView:
|
||||
self._server_i18n = server_i18n
|
||||
self._logger = logger
|
||||
|
||||
def _load_supporters(self) -> dict:
|
||||
def _load_supporters(self) -> dict[str, Any]:
|
||||
"""Load supporters data from JSON file."""
|
||||
try:
|
||||
current_file = os.path.abspath(__file__)
|
||||
@@ -152,7 +179,7 @@ class ModelPageView:
|
||||
self._template_env.filters["t"] = (
|
||||
self._server_i18n.create_template_filter()
|
||||
)
|
||||
self._template_env._i18n_filter_added = True # type: ignore[attr-defined]
|
||||
self._template_env._i18n_filter_added = True # pyright: ignore[reportAttributeAccessIssue]
|
||||
|
||||
from ...services.llm_service import PROVIDER_PRESETS
|
||||
|
||||
@@ -199,7 +226,7 @@ class ModelListingHandler:
|
||||
self,
|
||||
*,
|
||||
service,
|
||||
parse_specific_params: Callable[[web.Request], Dict],
|
||||
parse_specific_params: Callable[[web.Request], Dict[str, Any]],
|
||||
logger: logging.Logger,
|
||||
) -> None:
|
||||
self._service = service
|
||||
@@ -287,7 +314,7 @@ class ModelListingHandler:
|
||||
)
|
||||
return web.json_response({"error": str(exc)}, status=500)
|
||||
|
||||
def _parse_common_params(self, request: web.Request) -> Dict:
|
||||
def _parse_common_params(self, request: web.Request) -> Dict[str, Any]:
|
||||
page = int(request.query.get("page", "1"))
|
||||
page_size = min(int(request.query.get("page_size", "20")), 100)
|
||||
sort_by = request.query.get("sort_by", "name")
|
||||
@@ -341,6 +368,7 @@ class ModelListingHandler:
|
||||
== "true",
|
||||
"tags": request.query.get("search_tags", "false").lower() == "true",
|
||||
"creator": request.query.get("search_creator", "false").lower() == "true",
|
||||
"hash": request.query.get("search_hash", "false").lower() == "true",
|
||||
"recursive": request.query.get("recursive", "true").lower() == "true",
|
||||
}
|
||||
|
||||
@@ -394,12 +422,14 @@ class ModelListingHandler:
|
||||
)
|
||||
|
||||
# View-local-versions filter: show all local versions of a specific model
|
||||
# Accepts either a CivitAI modelId (int) or a HF group key like "hf:user/repo"
|
||||
civitai_model_id = request.query.get("civitai_model_id")
|
||||
if civitai_model_id is not None:
|
||||
try:
|
||||
civitai_model_id = int(civitai_model_id)
|
||||
except (TypeError, ValueError):
|
||||
civitai_model_id = None
|
||||
# Keep as string — could be an HF group key (e.g. "hf:user/repo")
|
||||
pass
|
||||
|
||||
return {
|
||||
"page": page,
|
||||
@@ -458,6 +488,7 @@ class ModelManagementHandler:
|
||||
return web.Response(text="Model path is required", status=400)
|
||||
|
||||
result = await self._lifecycle_service.delete_model(file_path)
|
||||
_broadcast_models_changed()
|
||||
return web.json_response(result)
|
||||
except ValueError as exc:
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=400)
|
||||
@@ -537,6 +568,7 @@ class ModelManagementHandler:
|
||||
# Update model_data with new hash
|
||||
model_data["sha256"] = sha256
|
||||
model_data["hash_status"] = "completed"
|
||||
hash_status = "completed"
|
||||
else:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "No SHA256 hash found"}, status=400
|
||||
@@ -544,6 +576,32 @@ class ModelManagementHandler:
|
||||
|
||||
await MetadataManager.hydrate_model_data(model_data)
|
||||
|
||||
# hydrate_model_data replaces model_data with .metadata.json content,
|
||||
# which may lack sha256. Restore from cache and persist the fix.
|
||||
if not model_data.get("sha256"):
|
||||
if sha256:
|
||||
model_data["sha256"] = sha256
|
||||
model_data["hash_status"] = model_data.get("hash_status", hash_status)
|
||||
data_to_save = model_data.copy()
|
||||
data_to_save.pop("folder", None)
|
||||
await MetadataManager.save_metadata(file_path, data_to_save)
|
||||
else:
|
||||
sha256 = await calculate_sha256(file_path)
|
||||
if sha256:
|
||||
model_data["sha256"] = sha256.lower()
|
||||
model_data["hash_status"] = "completed"
|
||||
data_to_save = model_data.copy()
|
||||
data_to_save.pop("folder", None)
|
||||
await MetadataManager.save_metadata(file_path, data_to_save)
|
||||
else:
|
||||
return web.json_response(
|
||||
{
|
||||
"success": False,
|
||||
"error": "Failed to compute SHA256 hash for model",
|
||||
},
|
||||
status=500,
|
||||
)
|
||||
|
||||
success, error = await self._metadata_sync.fetch_and_update_model(
|
||||
sha256=model_data["sha256"],
|
||||
file_path=file_path,
|
||||
@@ -566,7 +624,12 @@ class ModelManagementHandler:
|
||||
{"success": False, "error": OFFLINE_FRIENDLY_MESSAGE},
|
||||
status=503,
|
||||
)
|
||||
self._logger.error("Error fetching from CivitAI: %s", exc, exc_info=True)
|
||||
self._logger.error(
|
||||
"Error fetching from CivitAI for %s: %s",
|
||||
locals().get("file_path", "unknown"),
|
||||
exc,
|
||||
exc_info=True,
|
||||
)
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||
|
||||
async def relink_civitai(self, request: web.Request) -> web.Response:
|
||||
@@ -575,6 +638,16 @@ class ModelManagementHandler:
|
||||
file_path = data.get("file_path")
|
||||
model_id = data.get("model_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:
|
||||
return web.json_response(
|
||||
@@ -590,19 +663,32 @@ class ModelManagementHandler:
|
||||
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(
|
||||
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,
|
||||
**relink_kwargs
|
||||
)
|
||||
|
||||
await self._service.scanner.update_single_model_cache(
|
||||
file_path, file_path, updated_metadata
|
||||
)
|
||||
|
||||
message = f"Model successfully re-linked to Civitai model {model_id}" + (
|
||||
f" version {model_version_id}" if model_version_id else ""
|
||||
if source == "civarchive":
|
||||
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(
|
||||
{
|
||||
@@ -611,6 +697,8 @@ class ModelManagementHandler:
|
||||
"hash": updated_metadata.get("sha256", ""),
|
||||
}
|
||||
)
|
||||
except ValueError as exc:
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=400)
|
||||
except Exception as exc:
|
||||
if is_expected_offline_error(str(exc)):
|
||||
return web.json_response(
|
||||
@@ -624,7 +712,7 @@ class ModelManagementHandler:
|
||||
try:
|
||||
reader = await request.multipart()
|
||||
|
||||
field = await reader.next()
|
||||
field: Any = await reader.next()
|
||||
if field is None or field.name != "preview_file":
|
||||
raise ValueError("Expected 'preview_file' field")
|
||||
content_type = field.headers.get("Content-Type", "image/png")
|
||||
@@ -666,7 +754,7 @@ class ModelManagementHandler:
|
||||
{
|
||||
"success": True,
|
||||
"preview_url": config.get_preview_static_url(
|
||||
result["preview_path"]
|
||||
str(result["preview_path"])
|
||||
),
|
||||
"preview_nsfw_level": result["preview_nsfw_level"],
|
||||
}
|
||||
@@ -747,7 +835,7 @@ class ModelManagementHandler:
|
||||
|
||||
result = await self._preview_service.replace_preview(
|
||||
model_path=model_path,
|
||||
preview_data=preview_data,
|
||||
preview_data=preview_bytes,
|
||||
content_type=content_type,
|
||||
original_filename=original_filename,
|
||||
nsfw_level=nsfw_level,
|
||||
@@ -759,7 +847,7 @@ class ModelManagementHandler:
|
||||
{
|
||||
"success": True,
|
||||
"preview_url": config.get_preview_static_url(
|
||||
result["preview_path"]
|
||||
str(result["preview_path"])
|
||||
),
|
||||
"preview_nsfw_level": result["preview_nsfw_level"],
|
||||
}
|
||||
@@ -897,6 +985,8 @@ class ModelManagementHandler:
|
||||
file_path=file_path, new_file_name=new_file_name
|
||||
)
|
||||
|
||||
_broadcast_models_changed()
|
||||
|
||||
return web.json_response(
|
||||
{
|
||||
**result,
|
||||
@@ -925,6 +1015,7 @@ class ModelManagementHandler:
|
||||
)
|
||||
|
||||
result = await self._lifecycle_service.bulk_delete_models(file_paths)
|
||||
_broadcast_models_changed()
|
||||
return web.json_response(result)
|
||||
except ValueError as exc:
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=400)
|
||||
@@ -968,11 +1059,18 @@ class ModelQueryHandler:
|
||||
self._service = service
|
||||
self._logger = logger
|
||||
|
||||
@staticmethod
|
||||
def _parse_include_empty(request: web.Request) -> bool:
|
||||
"""Parse the include_empty query flag (``1``/``true``)."""
|
||||
return request.query.get("include_empty", "").lower() in ("1", "true")
|
||||
|
||||
async def get_top_tags(self, request: web.Request) -> web.Response:
|
||||
try:
|
||||
limit = int(request.query.get("limit", "20"))
|
||||
if limit < 0:
|
||||
limit = 20
|
||||
elif limit > 200:
|
||||
limit = 20
|
||||
top_tags = await self._service.get_top_tags(limit)
|
||||
return web.json_response({"success": True, "tags": top_tags})
|
||||
except Exception as exc:
|
||||
@@ -981,6 +1079,22 @@ class ModelQueryHandler:
|
||||
{"success": False, "error": "Internal server error"}, status=500
|
||||
)
|
||||
|
||||
async def search_tags(self, request: web.Request) -> web.Response:
|
||||
try:
|
||||
query = request.query.get("q", "")
|
||||
limit = int(request.query.get("limit", "20"))
|
||||
if limit < 0:
|
||||
limit = 20
|
||||
elif limit > 200:
|
||||
limit = 20
|
||||
tags = await self._service.search_tags(query, limit)
|
||||
return web.json_response({"success": True, "tags": tags})
|
||||
except Exception as exc:
|
||||
self._logger.error("Error searching tags: %s", exc, exc_info=True)
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Internal server error"}, status=500
|
||||
)
|
||||
|
||||
async def get_base_models(self, request: web.Request) -> web.Response:
|
||||
try:
|
||||
limit = int(request.query.get("limit", "20"))
|
||||
@@ -1009,6 +1123,7 @@ class ModelQueryHandler:
|
||||
await self._service.scan_models(
|
||||
force_refresh=True, rebuild_cache=full_rebuild
|
||||
)
|
||||
_broadcast_models_changed()
|
||||
if self._service.scanner.is_cancelled():
|
||||
return web.json_response(
|
||||
{
|
||||
@@ -1043,8 +1158,14 @@ class ModelQueryHandler:
|
||||
|
||||
async def get_folders(self, request: web.Request) -> web.Response:
|
||||
try:
|
||||
include_empty = self._parse_include_empty(request)
|
||||
if include_empty:
|
||||
# Live enumeration includes empty OS-created directories.
|
||||
folders = await self._service.scanner.get_all_folders()
|
||||
else:
|
||||
cache = await self._service.scanner.get_cached_data()
|
||||
return web.json_response({"folders": cache.folders})
|
||||
folders = cache.folders
|
||||
return web.json_response({"folders": folders})
|
||||
except Exception as exc:
|
||||
self._logger.error("Error getting folders: %s", exc)
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||
@@ -1069,7 +1190,9 @@ class ModelQueryHandler:
|
||||
{"success": False, "error": "model_root parameter is required"},
|
||||
status=400,
|
||||
)
|
||||
folder_tree = await self._service.get_folder_tree(model_root)
|
||||
folder_tree = await self._service.get_folder_tree(
|
||||
model_root, include_empty=self._parse_include_empty(request)
|
||||
)
|
||||
return web.json_response({"success": True, "tree": folder_tree})
|
||||
except Exception as exc:
|
||||
self._logger.error("Error getting folder tree: %s", exc)
|
||||
@@ -1077,7 +1200,9 @@ class ModelQueryHandler:
|
||||
|
||||
async def get_unified_folder_tree(self, request: web.Request) -> web.Response:
|
||||
try:
|
||||
unified_tree = await self._service.get_unified_folder_tree()
|
||||
unified_tree = await self._service.get_unified_folder_tree(
|
||||
include_empty=self._parse_include_empty(request)
|
||||
)
|
||||
return web.json_response({"success": True, "tree": unified_tree})
|
||||
except Exception as exc:
|
||||
self._logger.error("Error getting unified folder tree: %s", exc)
|
||||
@@ -1275,9 +1400,13 @@ class ModelQueryHandler:
|
||||
text=f"{self._service.model_type.capitalize()} file name is required",
|
||||
status=400,
|
||||
)
|
||||
notes = await self._service.get_model_notes(model_name)
|
||||
if notes is not None:
|
||||
return web.json_response({"success": True, "notes": notes})
|
||||
result = await self._service.get_model_notes(model_name)
|
||||
if result is not None:
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"notes": result["notes"],
|
||||
"file_path": result["file_path"],
|
||||
})
|
||||
return web.json_response(
|
||||
{
|
||||
"success": False,
|
||||
@@ -1432,8 +1561,111 @@ class ModelQueryHandler:
|
||||
search = request.query.get("search", "").strip()
|
||||
limit = min(int(request.query.get("limit", "15")), 100)
|
||||
offset = max(0, int(request.query.get("offset", "0")))
|
||||
|
||||
folder = request.query.get("folder")
|
||||
recursive = request.query.get("recursive", "true").lower() == "true"
|
||||
base_models = list(request.query.getall("base_model", []))
|
||||
model_types = list(request.query.getall("model_type", []))
|
||||
|
||||
tag_filters: Dict[str, str] = {}
|
||||
for tag in request.query.getall("tag_include", []):
|
||||
if tag:
|
||||
tag_filters[tag] = "include"
|
||||
for tag in request.query.getall("tag_exclude", []):
|
||||
if tag:
|
||||
tag_filters[tag] = "exclude"
|
||||
|
||||
auto_tag_filters: Dict[str, str] = {}
|
||||
for tag in request.query.getall("auto_tag_include", []):
|
||||
if tag:
|
||||
auto_tag_filters[tag] = "include"
|
||||
for tag in request.query.getall("auto_tag_exclude", []):
|
||||
if tag:
|
||||
auto_tag_filters[tag] = "exclude"
|
||||
|
||||
tag_logic = request.query.get("tag_logic", "any").lower()
|
||||
if tag_logic not in ("any", "all"):
|
||||
tag_logic = "any"
|
||||
|
||||
credit_required = request.query.get("credit_required")
|
||||
if credit_required is not None:
|
||||
credit_required = credit_required.lower() not in ("false", "0", "")
|
||||
|
||||
allow_selling_generated_content = request.query.get(
|
||||
"allow_selling_generated_content"
|
||||
)
|
||||
if allow_selling_generated_content is not None:
|
||||
allow_selling_generated_content = (
|
||||
allow_selling_generated_content.lower() not in ("false", "0", "")
|
||||
)
|
||||
|
||||
# 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
|
||||
# widget when filter mode is on) signals that the filter pipeline
|
||||
# must run even when no concrete filter is set, so global settings
|
||||
# like show_only_sfw stay consistent with the list endpoint.
|
||||
apply_filters = (
|
||||
use_active_filters
|
||||
or "recursive" in request.query
|
||||
or folder is not None
|
||||
or bool(base_models)
|
||||
or bool(model_types)
|
||||
or bool(tag_filters)
|
||||
or bool(auto_tag_filters)
|
||||
or credit_required is not None
|
||||
or allow_selling_generated_content is not None
|
||||
)
|
||||
|
||||
matching_paths = await self._service.search_relative_paths(
|
||||
search, limit, offset
|
||||
search,
|
||||
limit,
|
||||
offset,
|
||||
folder=folder,
|
||||
recursive=recursive,
|
||||
base_models=base_models,
|
||||
model_types=model_types,
|
||||
tags=tag_filters,
|
||||
auto_tags=auto_tag_filters,
|
||||
tag_logic=tag_logic,
|
||||
credit_required=credit_required,
|
||||
allow_selling_generated_content=allow_selling_generated_content,
|
||||
apply_filters=apply_filters,
|
||||
)
|
||||
return web.json_response(
|
||||
{"success": True, "relative_paths": matching_paths}
|
||||
@@ -1444,6 +1676,50 @@ class ModelQueryHandler:
|
||||
)
|
||||
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:
|
||||
"""Coordinate downloads and progress reporting."""
|
||||
@@ -1510,7 +1786,8 @@ class ModelDownloadHandler:
|
||||
import json
|
||||
|
||||
try:
|
||||
data["file_params"] = json.loads(file_params_json)
|
||||
# Normalize falsy payloads (e.g. {}) to None (#1058)
|
||||
data["file_params"] = json.loads(file_params_json) or None
|
||||
except json.JSONDecodeError:
|
||||
self._logger.warning(
|
||||
"Invalid file_params JSON: %s", file_params_json
|
||||
@@ -1662,7 +1939,8 @@ class ModelDownloadHandler:
|
||||
|
||||
model_id = int(model_id_str) if model_id_str else None
|
||||
model_version_id = int(model_version_id_str) if model_version_id_str else None
|
||||
file_params = json.loads(file_params_json) if file_params_json else None
|
||||
# Normalize falsy payloads (e.g. {}) to None (#1058)
|
||||
file_params = (json.loads(file_params_json) if file_params_json else None) or None
|
||||
|
||||
service = await DownloadQueueService.get_instance()
|
||||
item = await service.add_to_queue(
|
||||
@@ -1737,8 +2015,18 @@ class ModelDownloadHandler:
|
||||
try:
|
||||
status_filter = request.query.get("status") or None
|
||||
service = await DownloadQueueService.get_instance()
|
||||
cleared = await service.clear_queue(status_filter=status_filter)
|
||||
return web.json_response({"success": True, "cleared": cleared})
|
||||
cleared_ids = await service.clear_queue(status_filter=status_filter)
|
||||
# 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:
|
||||
self._logger.error(
|
||||
"Error clearing download queue: %s", exc, exc_info=True
|
||||
@@ -1783,14 +2071,20 @@ class ModelDownloadHandler:
|
||||
|
||||
async def delete_download_history_item(self, request: web.Request) -> web.Response:
|
||||
try:
|
||||
item_id = int(request.query.get("id", "0"))
|
||||
if not item_id:
|
||||
download_id = request.query.get("download_id")
|
||||
id_str = request.query.get("id")
|
||||
item_id = int(id_str) if id_str else None
|
||||
|
||||
if not download_id and not item_id:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "id is required"}, status=400
|
||||
{"success": False, "error": "id or download_id is required"},
|
||||
status=400,
|
||||
)
|
||||
|
||||
service = await DownloadQueueService.get_instance()
|
||||
deleted = await service.delete_history_item(item_id)
|
||||
deleted = await service.delete_history_item(
|
||||
id=item_id, download_id=download_id
|
||||
)
|
||||
return web.json_response({"success": deleted})
|
||||
except Exception as exc:
|
||||
self._logger.error(
|
||||
@@ -1800,18 +2094,26 @@ class ModelDownloadHandler:
|
||||
|
||||
async def retry_download_from_history(self, request: web.Request) -> web.Response:
|
||||
try:
|
||||
item_id = int(request.query.get("id", "0"))
|
||||
if not item_id:
|
||||
download_id = request.query.get("download_id")
|
||||
id_str = request.query.get("id")
|
||||
item_id = int(id_str) if id_str else None
|
||||
|
||||
if not download_id and not item_id:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "id is required"}, status=400
|
||||
{"success": False, "error": "id or download_id is required"},
|
||||
status=400,
|
||||
)
|
||||
|
||||
service = await DownloadQueueService.get_instance()
|
||||
item = await service.retry_from_history(item_id)
|
||||
item = await service.retry_from_history(
|
||||
item_id=item_id, download_id=download_id
|
||||
)
|
||||
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(
|
||||
{"success": False, "error": "History item not found or not retryable"},
|
||||
status=404,
|
||||
{"success": False, "error": "History item not found or not retryable"}
|
||||
)
|
||||
return web.json_response({"success": True, "item": item})
|
||||
except Exception as exc:
|
||||
@@ -1862,8 +2164,12 @@ class ModelDownloadHandler:
|
||||
completed_at=completed_at,
|
||||
)
|
||||
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(
|
||||
{"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})
|
||||
except Exception as exc:
|
||||
@@ -1905,9 +2211,10 @@ class ModelDownloadHandler:
|
||||
service = await DownloadQueueService.get_instance()
|
||||
updated = await service.update_status(download_id, status)
|
||||
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(
|
||||
{"success": False, "error": "Download not found in queue"},
|
||||
status=404,
|
||||
{"success": False, "error": "Download not found in queue"}
|
||||
)
|
||||
return web.json_response({"success": True})
|
||||
except Exception as exc:
|
||||
@@ -1927,7 +2234,7 @@ class ModelCivitaiHandler:
|
||||
settings_service: SettingsManager,
|
||||
ws_manager: WebSocketManager,
|
||||
logger: logging.Logger,
|
||||
metadata_provider_factory: Callable[[], Awaitable],
|
||||
metadata_provider_factory: Callable[[], Awaitable[Any]],
|
||||
validate_model_type: Callable[[str], bool],
|
||||
expected_model_types: Callable[[], str],
|
||||
find_model_file: Callable[
|
||||
@@ -1992,7 +2299,7 @@ class ModelCivitaiHandler:
|
||||
downloaded_version_ids = set(
|
||||
await history_service.get_downloaded_version_ids(
|
||||
self._service.model_type,
|
||||
model_id,
|
||||
int(model_id),
|
||||
)
|
||||
)
|
||||
except Exception as exc: # pragma: no cover - defensive logging
|
||||
@@ -2026,6 +2333,19 @@ class ModelCivitaiHandler:
|
||||
else:
|
||||
version.pop("localPath", None)
|
||||
|
||||
# Per-file downloaded state so multi-file versions can show
|
||||
# which individual files are already in the library (#1058)
|
||||
local_entries: List[Any] = []
|
||||
if version_id is not None and cache:
|
||||
files_getter = getattr(cache, "get_files_by_version_id", None)
|
||||
if files_getter is not None:
|
||||
local_entries = files_getter(version_id)
|
||||
elif cache_entry is not None:
|
||||
local_entries = [cache_entry]
|
||||
version["downloadedFiles"] = self._match_downloaded_files(
|
||||
version, local_entries
|
||||
)
|
||||
|
||||
model_file = (
|
||||
self._find_model_file(version.get("files", []))
|
||||
if isinstance(version.get("files"), Iterable)
|
||||
@@ -2040,6 +2360,64 @@ class ModelCivitaiHandler:
|
||||
)
|
||||
return web.Response(status=500, text=str(exc))
|
||||
|
||||
@staticmethod
|
||||
def _match_downloaded_files(
|
||||
version: Mapping[str, Any], local_entries: List[Any]
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Map local library entries back to individual files of a version.
|
||||
|
||||
Matching follows rule D2 (#1058): SHA256 is authoritative when the
|
||||
local entry carries one; otherwise fall back to extension-less file
|
||||
name equality. Returns ``[{fileId, fileName, filePath}]``.
|
||||
"""
|
||||
files = version.get("files")
|
||||
if not isinstance(files, list) or not local_entries:
|
||||
return []
|
||||
|
||||
by_hash: Dict[str, Mapping[str, Any]] = {}
|
||||
by_name: Dict[str, Mapping[str, Any]] = {}
|
||||
for file_info in files:
|
||||
if not isinstance(file_info, Mapping):
|
||||
continue
|
||||
sha = str(
|
||||
(file_info.get("hashes") or {}).get("SHA256") or ""
|
||||
).strip().lower()
|
||||
if sha:
|
||||
by_hash.setdefault(sha, file_info)
|
||||
name = str(file_info.get("name") or "").strip()
|
||||
if name:
|
||||
by_name.setdefault(os.path.splitext(name)[0], file_info)
|
||||
|
||||
downloaded: List[Dict[str, Any]] = []
|
||||
seen_keys: set = set()
|
||||
for entry in local_entries:
|
||||
if not isinstance(entry, Mapping):
|
||||
continue
|
||||
matched: Optional[Mapping[str, Any]] = None
|
||||
local_hash = str(entry.get("sha256") or "").strip().lower()
|
||||
if local_hash:
|
||||
matched = by_hash.get(local_hash)
|
||||
if matched is None:
|
||||
local_name = str(entry.get("file_name") or "").strip()
|
||||
if local_name:
|
||||
matched = by_name.get(local_name)
|
||||
if matched is None:
|
||||
continue
|
||||
|
||||
file_id = matched.get("id")
|
||||
dedupe_key = file_id if file_id is not None else matched.get("name")
|
||||
if dedupe_key in seen_keys:
|
||||
continue
|
||||
seen_keys.add(dedupe_key)
|
||||
downloaded.append(
|
||||
{
|
||||
"fileId": file_id,
|
||||
"fileName": matched.get("name"),
|
||||
"filePath": entry.get("file_path"),
|
||||
}
|
||||
)
|
||||
return downloaded
|
||||
|
||||
async def get_civitai_model_by_version(self, request: web.Request) -> web.Response:
|
||||
try:
|
||||
model_version_id = request.match_info.get("modelVersionId")
|
||||
@@ -2102,6 +2480,8 @@ class ModelMoveHandler:
|
||||
result = await self._move_service.move_model(
|
||||
file_path, target_path, use_default_paths=use_default_paths
|
||||
)
|
||||
if result.get("success"):
|
||||
_broadcast_models_changed()
|
||||
status = 200 if result.get("success") else 500
|
||||
return web.json_response(result, status=status)
|
||||
except Exception as exc:
|
||||
@@ -2121,6 +2501,8 @@ class ModelMoveHandler:
|
||||
result = await self._move_service.move_models_bulk(
|
||||
file_paths, target_path, use_default_paths=use_default_paths
|
||||
)
|
||||
if result.get("success"):
|
||||
_broadcast_models_changed()
|
||||
return web.json_response(result)
|
||||
except Exception as exc:
|
||||
self._logger.error("Error moving models in bulk: %s", exc, exc_info=True)
|
||||
@@ -2166,6 +2548,7 @@ class ModelAutoOrganizeHandler:
|
||||
progress_callback=self._progress_callback,
|
||||
exclusion_patterns=exclusion_patterns,
|
||||
)
|
||||
_broadcast_models_changed()
|
||||
return web.json_response(result.to_dict())
|
||||
except AutoOrganizeInProgressError:
|
||||
return web.json_response(
|
||||
@@ -2269,8 +2652,8 @@ class ModelUpdateHandler:
|
||||
self._logger.error("Failed to fetch license info: %s", exc, exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||
|
||||
updated: List[Dict[str, str]] = []
|
||||
errors: List[Dict[str, str]] = []
|
||||
updated: List[Dict[str, Any]] = []
|
||||
errors: List[Dict[str, Any]] = []
|
||||
for model_id in model_ids:
|
||||
license_payload = license_map.get(model_id)
|
||||
if not license_payload:
|
||||
@@ -2283,6 +2666,7 @@ class ModelUpdateHandler:
|
||||
model_section = civitai_section.get("model")
|
||||
if not isinstance(model_section, Mapping):
|
||||
model_section = {}
|
||||
model_section = dict(model_section)
|
||||
model_section.update(resolved_payload)
|
||||
civitai_section["model"] = model_section
|
||||
metadata_payload["civitai"] = civitai_section
|
||||
@@ -2298,7 +2682,7 @@ class ModelUpdateHandler:
|
||||
)
|
||||
errors.append({"filePath": metadata_path, "error": str(exc)})
|
||||
|
||||
response_payload = {"success": True, "updated": updated}
|
||||
response_payload: Dict[str, Any] = {"success": True, "updated": updated}
|
||||
missing_model_ids = [mid for mid in model_ids if mid not in license_map]
|
||||
if missing_model_ids:
|
||||
response_payload["missingModelIds"] = missing_model_ids
|
||||
@@ -2368,6 +2752,7 @@ class ModelUpdateHandler:
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||
|
||||
hide_early_access = False
|
||||
hide_paid = False
|
||||
if self._settings is not None:
|
||||
try:
|
||||
hide_early_access = bool(
|
||||
@@ -2375,12 +2760,27 @@ class ModelUpdateHandler:
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
hide_paid = bool(self._settings.get("hide_paid_updates", False))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
same_base_scope = self._uses_same_base_update_scope()
|
||||
|
||||
serialized_records = []
|
||||
for record in records.values():
|
||||
has_update_fn = getattr(record, "has_update", None)
|
||||
if callable(has_update_fn) and has_update_fn(
|
||||
hide_early_access=hide_early_access
|
||||
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_paid=hide_paid,
|
||||
):
|
||||
serialized_records.append(self._serialize_record(record))
|
||||
|
||||
@@ -2391,6 +2791,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:
|
||||
payload = await self._read_json(request)
|
||||
model_id = self._normalize_model_id(payload.get("modelId"))
|
||||
@@ -2534,10 +2954,16 @@ class ModelUpdateHandler:
|
||||
if not record or not record.versions:
|
||||
return record
|
||||
|
||||
# Find versions that need enrichment
|
||||
# Find versions that need enrichment. Permanent paid versions are not
|
||||
# early access (mirror _is_early_access_active) and never carry an end
|
||||
# time, so skip them to avoid pointless per-version API calls.
|
||||
versions_needing_update = []
|
||||
for version in record.versions:
|
||||
if version.is_early_access and not version.early_access_ends_at:
|
||||
if (
|
||||
version.is_early_access
|
||||
and not version.early_access_ends_at
|
||||
and not getattr(version, "is_paid", False)
|
||||
):
|
||||
versions_needing_update.append(version)
|
||||
|
||||
if not versions_needing_update:
|
||||
@@ -2647,6 +3073,7 @@ class ModelUpdateHandler:
|
||||
civitai_payload = metadata_payload.get("civitai")
|
||||
if not isinstance(civitai_payload, Mapping):
|
||||
civitai_payload = {}
|
||||
civitai_payload = dict(civitai_payload)
|
||||
|
||||
model_payload = civitai_payload.get("model")
|
||||
if not isinstance(model_payload, Mapping):
|
||||
@@ -2691,7 +3118,7 @@ class ModelUpdateHandler:
|
||||
|
||||
return aggregated
|
||||
|
||||
def _extract_target_model_ids(self, payload: Dict) -> Optional[List[int]]:
|
||||
def _extract_target_model_ids(self, payload: Dict[str, Any]) -> Optional[List[int]]:
|
||||
if not isinstance(payload, Mapping):
|
||||
return None
|
||||
|
||||
@@ -2719,7 +3146,7 @@ class ModelUpdateHandler:
|
||||
return {}
|
||||
|
||||
to_dict = getattr(metadata, "to_dict", None)
|
||||
if callable(to_dict):
|
||||
if to_dict:
|
||||
try:
|
||||
return to_dict()
|
||||
except Exception:
|
||||
@@ -2730,7 +3157,7 @@ class ModelUpdateHandler:
|
||||
|
||||
return {}
|
||||
|
||||
async def _read_json(self, request: web.Request) -> Dict:
|
||||
async def _read_json(self, request: web.Request) -> Dict[str, Any]:
|
||||
if not request.can_read_body:
|
||||
return {}
|
||||
try:
|
||||
@@ -2762,10 +3189,11 @@ class ModelUpdateHandler:
|
||||
record,
|
||||
*,
|
||||
version_context: Optional[Dict[int, Dict[str, Any]]] = None,
|
||||
) -> Dict:
|
||||
) -> Dict[str, Any]:
|
||||
context = version_context or {}
|
||||
# Check user setting for hiding early access versions
|
||||
hide_early_access = False
|
||||
hide_paid = False
|
||||
if self._settings is not None:
|
||||
try:
|
||||
hide_early_access = bool(
|
||||
@@ -2773,6 +3201,10 @@ class ModelUpdateHandler:
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
hide_paid = bool(self._settings.get("hide_paid_updates", False))
|
||||
except Exception:
|
||||
pass
|
||||
return {
|
||||
"modelType": record.model_type,
|
||||
"modelId": record.model_id,
|
||||
@@ -2781,7 +3213,10 @@ class ModelUpdateHandler:
|
||||
"inLibraryVersionIds": record.in_library_version_ids,
|
||||
"lastCheckedAt": record.last_checked_at,
|
||||
"shouldIgnore": record.should_ignore_model,
|
||||
"hasUpdate": record.has_update(hide_early_access=hide_early_access),
|
||||
"hasUpdate": record.has_update(
|
||||
hide_early_access=hide_early_access,
|
||||
hide_paid=hide_paid,
|
||||
),
|
||||
"versions": [
|
||||
self._serialize_version(version, context.get(version.version_id))
|
||||
for version in record.versions
|
||||
@@ -2791,7 +3226,7 @@ class ModelUpdateHandler:
|
||||
@staticmethod
|
||||
def _serialize_version(
|
||||
version, context: Optional[Dict[str, Any]]
|
||||
) -> Dict:
|
||||
) -> Dict[str, Any]:
|
||||
context = context or {}
|
||||
preview_override = context.get("preview_override")
|
||||
preview_url = (
|
||||
@@ -2800,8 +3235,11 @@ class ModelUpdateHandler:
|
||||
|
||||
# Determine if version is currently in early access
|
||||
# Two-phase detection: use exact end time if available, otherwise fallback to basic flag
|
||||
# Mirror _is_early_access_active: permanent paid versions (no end time) are NOT early access
|
||||
is_early_access = False
|
||||
if version.early_access_ends_at:
|
||||
if getattr(version, "is_paid", False) and not version.early_access_ends_at:
|
||||
is_early_access = False
|
||||
elif version.early_access_ends_at:
|
||||
try:
|
||||
from datetime import datetime, timezone
|
||||
|
||||
@@ -2816,6 +3254,13 @@ class ModelUpdateHandler:
|
||||
# Fallback to basic EA flag from bulk API
|
||||
is_early_access = True
|
||||
|
||||
paid_access_payload = None
|
||||
if getattr(version, "paid_access", None):
|
||||
try:
|
||||
paid_access_payload = json.loads(version.paid_access)
|
||||
except (TypeError, ValueError):
|
||||
paid_access_payload = None
|
||||
|
||||
return {
|
||||
"versionId": version.version_id,
|
||||
"name": version.name,
|
||||
@@ -2829,8 +3274,13 @@ class ModelUpdateHandler:
|
||||
"earlyAccessEndsAt": version.early_access_ends_at,
|
||||
"isEarlyAccess": is_early_access,
|
||||
"usageControl": version.usage_control,
|
||||
"isPaid": bool(getattr(version, "is_paid", False)),
|
||||
"paidAccess": paid_access_payload,
|
||||
"filePath": context.get("file_path"),
|
||||
"fileName": context.get("file_name"),
|
||||
# Weight-file variant count (None when unknown); lets the UI hide
|
||||
# the download affordance for single-file in-library versions.
|
||||
"fileCount": getattr(version, "file_count", None),
|
||||
}
|
||||
|
||||
async def _build_version_context(
|
||||
@@ -2931,6 +3381,7 @@ class ModelHandlerSet:
|
||||
"bulk_delete_models": self.management.bulk_delete_models,
|
||||
"verify_duplicates": self.management.verify_duplicates,
|
||||
"get_top_tags": self.query.get_top_tags,
|
||||
"search_tags": self.query.search_tags,
|
||||
"get_base_models": self.query.get_base_models,
|
||||
"get_model_types": self.query.get_model_types,
|
||||
"scan_models": self.query.scan_models,
|
||||
@@ -2974,6 +3425,8 @@ class ModelHandlerSet:
|
||||
"get_model_metadata": self.query.get_model_metadata,
|
||||
"get_model_description": self.query.get_model_description,
|
||||
"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,
|
||||
"fetch_missing_civitai_license_data": self.updates.fetch_missing_civitai_license_data,
|
||||
"set_model_update_ignore": self.updates.set_model_update_ignore,
|
||||
|
||||
@@ -0,0 +1,323 @@
|
||||
"""Handler for the pending-delete undo endpoint.
|
||||
|
||||
Restores a staged delete batch (models or recipes) via
|
||||
``PendingDeleteService.undo`` and then repairs the affected library caches:
|
||||
the model cache entry is restored from the manifest's ``model_snapshot``
|
||||
(including the version index and hash index), tag counts are re-incremented,
|
||||
and the recipe cache is re-populated via ``RecipeScanner.add_recipe``.
|
||||
|
||||
The per-type scanner is resolved from the manifest's ``model_type`` page value
|
||||
through the SAME ServiceRegistry getters the model route registrars use
|
||||
(lora/checkpoint/embedding) - never a hardcoded lora scanner.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
from typing import Any, Awaitable, Callable, Dict, List, Optional, Set, cast
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
from ...services.pending_delete_service import get_pending_delete_service
|
||||
from .model_handlers import _broadcast_models_changed
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Manifest ``model_type`` page values -> ServiceRegistry scanner getter names.
|
||||
# The model route registrars resolve per-type scanners via these getters
|
||||
# (lora_routes / checkpoint_routes / embedding_routes); undo must do the same
|
||||
# so the CORRECT cache is restored for the deleted model's type.
|
||||
_MODEL_TYPE_GETTER_NAMES: Dict[str, str] = {
|
||||
"loras": "get_lora_scanner",
|
||||
"checkpoints": "get_checkpoint_scanner",
|
||||
"embeddings": "get_embedding_scanner",
|
||||
}
|
||||
|
||||
# Staged batch ids are ``uuid.uuid4().hex`` (32 lowercase hex chars). The id is
|
||||
# joined into filesystem paths by ``_find_batch_dir``, so reject anything that
|
||||
# does not match this exact shape (blocks path-traversal via batch_id).
|
||||
_BATCH_ID_RE = re.compile(r"^[0-9a-f]{32}$")
|
||||
|
||||
|
||||
class PendingDeleteHandler:
|
||||
"""Handle undo requests for staged model/recipe deletions."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
service_factory: Callable[[], Awaitable[Any]] = get_pending_delete_service,
|
||||
scanner_getter: Optional[Callable[[str], Awaitable[Any]]] = None,
|
||||
recipe_scanner_getter: Optional[Callable[[], Awaitable[Any]]] = None,
|
||||
) -> None:
|
||||
self._service_factory: Callable[[], Awaitable[Any]] = service_factory
|
||||
self._scanner_getter: Callable[[str], Awaitable[Any]] = (
|
||||
scanner_getter or self._resolve_scanner
|
||||
)
|
||||
self._recipe_scanner_getter: Callable[[], Awaitable[Any]] = (
|
||||
recipe_scanner_getter or self._resolve_recipe_scanner
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _resolve_scanner(model_type: str) -> Any:
|
||||
"""Resolve the per-type scanner for a manifest ``model_type``.
|
||||
|
||||
The getter is looked up on the ServiceRegistry module namespace at call
|
||||
time so tests (and the registry stubs) can patch it.
|
||||
"""
|
||||
from ...services import service_registry
|
||||
|
||||
getter_name = _MODEL_TYPE_GETTER_NAMES.get(model_type)
|
||||
if getter_name is None:
|
||||
raise ValueError(f"Unknown model type: {model_type}")
|
||||
getter = getattr(service_registry.ServiceRegistry, getter_name, None)
|
||||
if not callable(getter):
|
||||
raise ValueError(f"No scanner getter for model type: {model_type}")
|
||||
scanner = await cast(Callable[[], Awaitable[Any]], getter)()
|
||||
if scanner is None:
|
||||
raise ValueError(f"No scanner registered for model type: {model_type}")
|
||||
return scanner
|
||||
|
||||
@staticmethod
|
||||
async def _resolve_recipe_scanner() -> Any:
|
||||
"""Resolve the recipe scanner via the ServiceRegistry module namespace."""
|
||||
from ...services import service_registry
|
||||
|
||||
getter = getattr(service_registry.ServiceRegistry, "get_recipe_scanner", None)
|
||||
if not callable(getter):
|
||||
raise ValueError("Recipe scanner getter unavailable")
|
||||
scanner = await cast(Callable[[], Awaitable[Any]], getter)()
|
||||
if scanner is None:
|
||||
raise ValueError("No recipe scanner registered")
|
||||
return scanner
|
||||
|
||||
async def undo_delete(self, request: web.Request) -> web.Response:
|
||||
"""Restore a staged batch and its library cache entry.
|
||||
|
||||
Body: ``{"batch_id": str}``. On success returns
|
||||
``{"success": True, "restored": [<original paths>], "kind": kind}``.
|
||||
Expired/unknown batches and occupied target paths -> 404.
|
||||
"""
|
||||
try:
|
||||
data = await request.json()
|
||||
except Exception:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid JSON body"}, status=400
|
||||
)
|
||||
if not isinstance(data, dict):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid JSON body"}, status=400
|
||||
)
|
||||
batch_id = data.get("batch_id")
|
||||
if not batch_id or not isinstance(batch_id, str):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "batch_id is required"}, status=400
|
||||
)
|
||||
if not _BATCH_ID_RE.fullmatch(batch_id):
|
||||
# batch_id is joined into a path by _find_batch_dir - restrict to
|
||||
# the exact staged-id shape so traversal payloads get 400.
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid batch_id"}, status=400
|
||||
)
|
||||
|
||||
service = await self._service_factory()
|
||||
try:
|
||||
# Read the manifest BEFORE undo: undo() removes the batch dir.
|
||||
manifest = await self._read_staged_manifest(service, batch_id)
|
||||
result = await service.undo(batch_id)
|
||||
except ValueError as exc:
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=404)
|
||||
except Exception as exc:
|
||||
logger.error("Unexpected error undoing batch %s: %s", batch_id, exc, exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||
|
||||
kind = result.get("kind")
|
||||
try:
|
||||
if kind == "model":
|
||||
if manifest is not None:
|
||||
await self._restore_model_cache(manifest)
|
||||
else:
|
||||
# undo() raises when the manifest is missing, so this only
|
||||
# happens defensively - files are restored regardless.
|
||||
logger.warning(
|
||||
"Manifest missing after undo of %s; skipping cache restore",
|
||||
batch_id,
|
||||
)
|
||||
_broadcast_models_changed()
|
||||
elif kind == "recipe":
|
||||
# Recipe undo is client-refresh only: re-add to the scanner
|
||||
# cache, no models_changed broadcast.
|
||||
if manifest is not None:
|
||||
await self._restore_recipe_cache(result, manifest)
|
||||
else:
|
||||
logger.warning(
|
||||
"Manifest missing after undo of %s; skipping cache restore",
|
||||
batch_id,
|
||||
)
|
||||
except Exception as exc:
|
||||
# Files are already restored; only the cache restoration failed.
|
||||
logger.error(
|
||||
"Cache restoration failed after undo of %s: %s",
|
||||
batch_id,
|
||||
exc,
|
||||
exc_info=True,
|
||||
)
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||
|
||||
return web.json_response(
|
||||
{
|
||||
"success": True,
|
||||
"restored": result.get("restored", []),
|
||||
"kind": kind,
|
||||
}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _read_staged_manifest(
|
||||
service: Any, batch_id: str
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Locate and read the batch manifest while it still exists on disk."""
|
||||
batch_dir = await service._find_batch_dir(batch_id)
|
||||
if not batch_dir:
|
||||
return None
|
||||
manifest_path = os.path.join(batch_dir, "manifest.json")
|
||||
try:
|
||||
with open(manifest_path, "r", encoding="utf-8") as handle:
|
||||
payload = json.load(handle)
|
||||
except (OSError, json.JSONDecodeError) as exc:
|
||||
logger.debug("Failed to read manifest for batch %s: %s", batch_id, exc)
|
||||
return None
|
||||
return payload if isinstance(payload, dict) else None
|
||||
|
||||
async def _restore_model_cache(self, manifest: Dict[str, Any]) -> None:
|
||||
"""Re-add every deleted model's cache entry from the manifest.
|
||||
|
||||
Each main-file entry carries the deleted model's ``snapshot`` (added at
|
||||
stage time), so a merged bulk manifest holds ALL snapshots - undo must
|
||||
restore every one, not just the top-level winner's. Old-format
|
||||
manifests without entry snapshots fall back to the top-level
|
||||
``model_snapshot`` (backward compat / single-delete path).
|
||||
"""
|
||||
model_type = manifest.get("model_type")
|
||||
if not model_type or not isinstance(model_type, str):
|
||||
raise ValueError(f"Manifest carries no model_type: {manifest.get('batch_id')}")
|
||||
scanner = await self._scanner_getter(model_type)
|
||||
|
||||
# Collect one snapshot per distinct file_path from the entry snapshots.
|
||||
snapshots: List[Dict[str, Any]] = []
|
||||
seen: Set[str] = set()
|
||||
for entry in manifest.get("entries") or []:
|
||||
snapshot = entry.get("snapshot")
|
||||
if not isinstance(snapshot, dict):
|
||||
continue
|
||||
file_path = snapshot.get("file_path")
|
||||
if not file_path or not isinstance(file_path, str):
|
||||
continue
|
||||
if file_path in seen:
|
||||
continue
|
||||
seen.add(file_path)
|
||||
snapshots.append(snapshot)
|
||||
|
||||
if not snapshots:
|
||||
# Backward compat: pre-F3 manifests carry only the top-level
|
||||
# model_snapshot (single-delete path, unchanged behavior).
|
||||
top = manifest.get("model_snapshot")
|
||||
if isinstance(top, dict) and top.get("file_path"):
|
||||
snapshots = [top]
|
||||
else:
|
||||
logger.warning(
|
||||
"Manifest %s has no restorable model snapshot; skipping cache restore",
|
||||
manifest.get("batch_id"),
|
||||
)
|
||||
return
|
||||
|
||||
cache = await scanner.get_cached_data()
|
||||
if cache is None:
|
||||
logger.warning(
|
||||
"Scanner cache unavailable for %s; skipping cache restore", model_type
|
||||
)
|
||||
return
|
||||
|
||||
for snapshot in snapshots:
|
||||
file_path = str(snapshot["file_path"])
|
||||
# A rescan between delete and undo may have re-added a stale entry
|
||||
# for this path - drop it so exactly one (the snapshot) remains.
|
||||
cache.raw_data = [
|
||||
item for item in cache.raw_data if item.get("file_path") != file_path
|
||||
]
|
||||
|
||||
# Restore tag counts (mirror of the bulk-delete decrement in
|
||||
# _batch_update_cache_for_deleted_models: undo re-increments).
|
||||
tags = snapshot.get("tags")
|
||||
if isinstance(tags, list):
|
||||
for tag in tags:
|
||||
if not isinstance(tag, str) or not tag:
|
||||
continue
|
||||
scanner._tags_count[tag] = scanner._tags_count.get(tag, 0) + 1
|
||||
|
||||
cache.raw_data.append(dict(snapshot))
|
||||
|
||||
# Re-register the path in the hash index (add_entry guards a
|
||||
# missing sha256 internally; still guard defensively here).
|
||||
sha256 = snapshot.get("sha256") or ""
|
||||
autov3 = snapshot.get("autov3")
|
||||
hash_index = getattr(scanner, "_hash_index", None)
|
||||
if hash_index is not None and sha256 and file_path:
|
||||
hash_index.add_entry(sha256, file_path, autov3)
|
||||
|
||||
# Follow the bulk-delete cache-update pattern ONCE after all entries,
|
||||
# including the explicit version-index rebuild so the version index
|
||||
# does not go stale.
|
||||
cache.rebuild_version_index()
|
||||
await cache.resort()
|
||||
|
||||
scanner.bump_cache_version()
|
||||
|
||||
persist = getattr(scanner, "_persist_current_cache", None)
|
||||
if callable(persist):
|
||||
result = persist()
|
||||
if inspect.isawaitable(result):
|
||||
await result
|
||||
|
||||
async def _restore_recipe_cache(
|
||||
self, result: Dict[str, Any], manifest: Dict[str, Any]
|
||||
) -> None:
|
||||
"""Re-add a restored recipe via ``RecipeScanner.add_recipe``.
|
||||
|
||||
The recipe JSON embeds the full recipe_data (incl. id/file_path);
|
||||
``add_recipe`` only READS the ``_json_path_map`` so the forced frontend
|
||||
refresh self-heals any transient path-map gap.
|
||||
"""
|
||||
restored = result.get("restored") or []
|
||||
json_path = next(
|
||||
(p for p in restored if isinstance(p, str) and p.endswith(".json")),
|
||||
None,
|
||||
)
|
||||
if not json_path or not os.path.exists(json_path):
|
||||
# Defensive fallback to the manifest's recipe_snapshot file_path.
|
||||
snapshot = manifest.get("recipe_snapshot") or {}
|
||||
fallback = snapshot.get("file_path")
|
||||
if fallback and os.path.exists(fallback):
|
||||
json_path = fallback
|
||||
else:
|
||||
logger.warning(
|
||||
"Restored recipe JSON not found in %s; skipping cache restore",
|
||||
restored,
|
||||
)
|
||||
return
|
||||
try:
|
||||
with open(json_path, "r", encoding="utf-8") as handle:
|
||||
recipe_data = json.load(handle)
|
||||
except (OSError, json.JSONDecodeError) as exc:
|
||||
logger.warning("Failed to load restored recipe JSON %s: %s", json_path, exc)
|
||||
return
|
||||
if not isinstance(recipe_data, dict):
|
||||
return
|
||||
recipe_scanner = await self._recipe_scanner_getter()
|
||||
await recipe_scanner.add_recipe(recipe_data)
|
||||
|
||||
|
||||
__all__ = ["PendingDeleteHandler"]
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,8 +1,8 @@
|
||||
import asyncio
|
||||
import logging
|
||||
from aiohttp import web
|
||||
from typing import Dict
|
||||
from server import PromptServer # type: ignore
|
||||
from typing import Any, Dict
|
||||
from server import PromptServer # pyright: ignore[reportMissingImports]
|
||||
|
||||
from .base_model_routes import BaseModelRoutes
|
||||
from .model_route_registrar import ModelRouteRegistrar
|
||||
@@ -31,13 +31,13 @@ class LoraRoutes(BaseModelRoutes):
|
||||
# Attach service dependencies
|
||||
self.attach_service(self.service)
|
||||
|
||||
def setup_routes(self, app: web.Application):
|
||||
def setup_routes(self, app: web.Application, prefix: str = "loras"):
|
||||
"""Setup LoRA routes"""
|
||||
# Schedule service initialization on app startup
|
||||
app.on_startup.append(lambda _: self.initialize_services())
|
||||
|
||||
# Setup common routes with 'loras' prefix (includes page route)
|
||||
super().setup_routes(app, "loras")
|
||||
super().setup_routes(app, prefix)
|
||||
|
||||
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
|
||||
"""Setup LoRA-specific routes"""
|
||||
@@ -73,7 +73,7 @@ class LoraRoutes(BaseModelRoutes):
|
||||
"POST", "/api/lm/{prefix}/get_trigger_words", prefix, self.get_trigger_words
|
||||
)
|
||||
|
||||
def _parse_specific_params(self, request: web.Request) -> Dict:
|
||||
def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]:
|
||||
"""Parse LoRA-specific parameters"""
|
||||
params = {}
|
||||
|
||||
@@ -119,25 +119,6 @@ class LoraRoutes(BaseModelRoutes):
|
||||
logger.error(f"Error getting letter counts: {e}")
|
||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
||||
|
||||
async def get_lora_notes(self, request: web.Request) -> web.Response:
|
||||
"""Get notes for a specific LoRA file"""
|
||||
try:
|
||||
lora_name = request.query.get("name")
|
||||
if not lora_name:
|
||||
return web.Response(text="Lora file name is required", status=400)
|
||||
|
||||
notes = await self.service.get_lora_notes(lora_name)
|
||||
if notes is not None:
|
||||
return web.json_response({"success": True, "notes": notes})
|
||||
else:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "LoRA not found in cache"}, status=404
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting lora notes: {e}", exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
||||
|
||||
async def get_lora_trigger_words(self, request: web.Request) -> web.Response:
|
||||
"""Get trigger words for a specific LoRA file"""
|
||||
try:
|
||||
@@ -168,52 +149,6 @@ class LoraRoutes(BaseModelRoutes):
|
||||
logger.error(f"Error getting lora usage tips by path: {e}", exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
||||
|
||||
async def get_lora_preview_url(self, request: web.Request) -> web.Response:
|
||||
"""Get the static preview URL for a LoRA file"""
|
||||
try:
|
||||
lora_name = request.query.get("name")
|
||||
if not lora_name:
|
||||
return web.Response(text="Lora file name is required", status=400)
|
||||
|
||||
preview_url = await self.service.get_lora_preview_url(lora_name)
|
||||
if preview_url:
|
||||
return web.json_response({"success": True, "preview_url": preview_url})
|
||||
else:
|
||||
return web.json_response(
|
||||
{
|
||||
"success": False,
|
||||
"error": "No preview URL found for the specified lora",
|
||||
},
|
||||
status=404,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting lora preview URL: {e}", exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
||||
|
||||
async def get_lora_civitai_url(self, request: web.Request) -> web.Response:
|
||||
"""Get the Civitai URL for a LoRA file"""
|
||||
try:
|
||||
lora_name = request.query.get("name")
|
||||
if not lora_name:
|
||||
return web.Response(text="Lora file name is required", status=400)
|
||||
|
||||
result = await self.service.get_lora_civitai_url(lora_name)
|
||||
if result["civitai_url"]:
|
||||
return web.json_response({"success": True, **result})
|
||||
else:
|
||||
return web.json_response(
|
||||
{
|
||||
"success": False,
|
||||
"error": "No Civitai data found for the specified lora",
|
||||
},
|
||||
status=404,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting lora Civitai URL: {e}", exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
||||
|
||||
async def get_random_loras(self, request: web.Request) -> web.Response:
|
||||
"""Get random LoRAs based on filters and strength ranges"""
|
||||
try:
|
||||
@@ -337,7 +272,7 @@ class LoraRoutes(BaseModelRoutes):
|
||||
graph_identifier = entry.get("graph_id")
|
||||
|
||||
try:
|
||||
parsed_node_id = int(node_identifier)
|
||||
parsed_node_id = int(node_identifier) # pyright: ignore[reportArgumentType]
|
||||
except (TypeError, ValueError):
|
||||
parsed_node_id = node_identifier
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ miscellaneous endpoints share a consistent registration flow.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Iterable, Mapping
|
||||
from typing import Any, Callable, Iterable, Mapping
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
@@ -32,6 +32,7 @@ MISC_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
||||
RouteDefinition("GET", "/api/lm/settings/libraries", "get_settings_libraries"),
|
||||
RouteDefinition("POST", "/api/lm/settings/libraries/activate", "activate_library"),
|
||||
RouteDefinition("GET", "/api/lm/health-check", "health_check"),
|
||||
RouteDefinition("GET", "/api/lm/init-status", "get_init_status"),
|
||||
RouteDefinition("GET", "/api/lm/supporters", "get_supporters"),
|
||||
RouteDefinition("GET", "/api/lm/wildcards/search", "search_wildcards"),
|
||||
RouteDefinition("POST", "/api/lm/wildcards/open-location", "open_wildcards_location"),
|
||||
@@ -39,10 +40,12 @@ MISC_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
||||
RouteDefinition("POST", "/api/lm/update-usage-stats", "update_usage_stats"),
|
||||
RouteDefinition("GET", "/api/lm/get-usage-stats", "get_usage_stats"),
|
||||
RouteDefinition("POST", "/api/lm/update-lora-code", "update_lora_code"),
|
||||
RouteDefinition("GET", "/api/lm/update-lora-code", "get_update_lora_code"),
|
||||
RouteDefinition("GET", "/api/lm/trained-words", "get_trained_words"),
|
||||
RouteDefinition("GET", "/api/lm/model-example-files", "get_model_example_files"),
|
||||
RouteDefinition("POST", "/api/lm/register-nodes", "register_nodes"),
|
||||
RouteDefinition("POST", "/api/lm/update-node-widget", "update_node_widget"),
|
||||
RouteDefinition("GET", "/api/lm/update-node-widget", "get_update_node_widget"),
|
||||
RouteDefinition("GET", "/api/lm/get-registry", "get_registry"),
|
||||
RouteDefinition("GET", "/api/lm/check-model-exists", "check_model_exists"),
|
||||
RouteDefinition("GET", "/api/lm/check-models-exist", "check_models_exist"),
|
||||
@@ -145,7 +148,7 @@ class MiscRouteRegistrar:
|
||||
handler_lookup[definition.handler_name],
|
||||
)
|
||||
|
||||
def _bind(self, method: str, path: str, handler: Callable) -> None:
|
||||
def _bind(self, method: str, path: str, handler: Callable[..., Any]) -> None:
|
||||
add_method_name = self._METHOD_MAP[method.upper()]
|
||||
add_method = getattr(self._app.router, add_method_name)
|
||||
add_method(path, handler)
|
||||
|
||||
@@ -7,7 +7,7 @@ import os
|
||||
from typing import Awaitable, Callable, Mapping
|
||||
|
||||
from aiohttp import web
|
||||
from server import PromptServer # type: ignore
|
||||
from server import PromptServer # pyright: ignore[reportMissingImports]
|
||||
|
||||
from ..services.metadata_service import (
|
||||
get_metadata_archive_manager,
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Iterable, Mapping
|
||||
from typing import Any, Callable, Iterable, Mapping
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
@@ -46,6 +46,7 @@ COMMON_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
||||
"GET", "/api/lm/{prefix}/auto-organize-progress", "get_auto_organize_progress"
|
||||
),
|
||||
RouteDefinition("GET", "/api/lm/{prefix}/top-tags", "get_top_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}/model-types", "get_model_types"),
|
||||
RouteDefinition("GET", "/api/lm/{prefix}/scan", "scan_models"),
|
||||
@@ -67,6 +68,8 @@ COMMON_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
||||
"GET", "/api/lm/{prefix}/model-description", "get_model_description"
|
||||
),
|
||||
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(
|
||||
"GET", "/api/lm/{prefix}/civitai/versions/{model_id}", "get_civitai_versions"
|
||||
),
|
||||
@@ -173,15 +176,15 @@ class ModelRouteRegistrar:
|
||||
handler_lookup[definition.handler_name],
|
||||
)
|
||||
|
||||
def add_route(self, method: str, path: str, handler: Callable) -> None:
|
||||
def add_route(self, method: str, path: str, handler: Callable[..., Any]) -> None:
|
||||
self._bind_route(method, path, handler)
|
||||
|
||||
def add_prefixed_route(
|
||||
self, method: str, path_template: str, prefix: str, handler: Callable
|
||||
self, method: str, path_template: str, prefix: str, handler: Callable[..., Any]
|
||||
) -> None:
|
||||
self._bind_route(method, path_template.replace("{prefix}", prefix), handler)
|
||||
|
||||
def _bind_route(self, method: str, path: str, handler: Callable) -> None:
|
||||
def _bind_route(self, method: str, path: str, handler: Callable[..., Any]) -> None:
|
||||
add_method_name = self._METHOD_MAP[method.upper()]
|
||||
add_method = getattr(self._app.router, add_method_name)
|
||||
add_method(path, handler)
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
"""Route controller for the pending-delete undo endpoint."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
from .handlers.pending_delete_handler import PendingDeleteHandler
|
||||
|
||||
|
||||
class PendingDeleteRoutes:
|
||||
"""Shared route controller mirroring MiscRoutes/UpdateRoutes.
|
||||
|
||||
Registered ONCE per mode (py/lora_manager.py, standalone.py); NEVER through
|
||||
the per-model-type ModelRouteRegistrar, which is instantiated per model
|
||||
type and would register this non-prefixed route three times.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def setup_routes(app: web.Application) -> None:
|
||||
"""Register the shared undo-delete endpoint."""
|
||||
handler = PendingDeleteHandler()
|
||||
_ = app.router.add_post("/api/lm/undo-delete", handler.undo_delete)
|
||||
|
||||
|
||||
__all__ = ["PendingDeleteRoutes"]
|
||||
@@ -3,7 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Mapping
|
||||
from typing import Any, Callable, Mapping
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
@@ -29,6 +29,7 @@ ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
||||
RouteDefinition("POST", "/api/lm/recipes/save", "save_recipe"),
|
||||
RouteDefinition("DELETE", "/api/lm/recipe/{recipe_id}", "delete_recipe"),
|
||||
RouteDefinition("GET", "/api/lm/recipes/top-tags", "get_top_tags"),
|
||||
RouteDefinition("GET", "/api/lm/recipes/search-tags", "search_tags"),
|
||||
RouteDefinition("GET", "/api/lm/recipes/base-models", "get_base_models"),
|
||||
RouteDefinition("GET", "/api/lm/recipes/roots", "get_roots"),
|
||||
RouteDefinition("GET", "/api/lm/recipes/folders", "get_folders"),
|
||||
@@ -42,9 +43,37 @@ ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
||||
),
|
||||
RouteDefinition("GET", "/api/lm/recipe/{recipe_id}/syntax", "get_recipe_syntax"),
|
||||
RouteDefinition("PUT", "/api/lm/recipe/{recipe_id}/update", "update_recipe"),
|
||||
RouteDefinition(
|
||||
"POST", "/api/lm/recipe/{recipe_id}/opened", "record_recipe_open"
|
||||
),
|
||||
RouteDefinition("POST", "/api/lm/recipe/move", "move_recipe"),
|
||||
RouteDefinition("POST", "/api/lm/recipes/move-bulk", "move_recipes_bulk"),
|
||||
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("POST", "/api/lm/recipes/bulk-delete", "bulk_delete"),
|
||||
RouteDefinition(
|
||||
@@ -55,11 +84,11 @@ ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
||||
"GET", "/api/lm/recipes/for-checkpoint", "get_recipes_for_checkpoint"
|
||||
),
|
||||
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-bulk", "rematch_recipes_bulk"),
|
||||
RouteDefinition("POST", "/api/lm/recipe/{recipe_id}/rematch", "rematch_recipe"),
|
||||
RouteDefinition("POST", "/api/lm/recipes/cancel-rematch", "cancel_rematch"),
|
||||
RouteDefinition("GET", "/api/lm/recipes/rematch-progress", "get_rematch_progress"),
|
||||
RouteDefinition("POST", "/api/lm/recipes/batch-import/start", "start_batch_import"),
|
||||
RouteDefinition(
|
||||
"GET", "/api/lm/recipes/batch-import/progress", "get_batch_import_progress"
|
||||
@@ -81,6 +110,14 @@ ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
||||
RouteDefinition(
|
||||
"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(
|
||||
"POST", "/api/lm/recipe/{recipe_id}/send-workflow", "send_recipe_workflow"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -104,7 +141,7 @@ class RecipeRouteRegistrar:
|
||||
handler = handler_lookup[definition.handler_name]
|
||||
self._bind_route(definition.method, definition.path, handler)
|
||||
|
||||
def _bind_route(self, method: str, path: str, handler: Callable) -> None:
|
||||
def _bind_route(self, method: str, path: str, handler: Callable[..., Any]) -> None:
|
||||
add_method_name = self._METHOD_MAP[method.upper()]
|
||||
add_method = getattr(self._app.router, add_method_name)
|
||||
add_method(path, handler)
|
||||
|
||||
+11
-10
@@ -40,10 +40,11 @@ class StatsRoutes:
|
||||
"""Route handlers for Statistics page and API endpoints"""
|
||||
|
||||
def __init__(self):
|
||||
self.lora_scanner = None
|
||||
self.checkpoint_scanner = None
|
||||
self.embedding_scanner = None
|
||||
self.usage_stats = None
|
||||
self.lora_scanner: Any = None
|
||||
self.checkpoint_scanner: Any = None
|
||||
self.embedding_scanner: Any = None
|
||||
self.usage_stats: Any = None
|
||||
self._i18n_filter_added = False
|
||||
self.template_env = jinja2.Environment(
|
||||
loader=jinja2.FileSystemLoader(config.templates_path),
|
||||
autoescape=True
|
||||
@@ -95,9 +96,9 @@ class StatsRoutes:
|
||||
server_i18n.set_locale(user_language)
|
||||
|
||||
# 为模板环境添加i18n过滤器
|
||||
if not hasattr(self.template_env, '_i18n_filter_added'):
|
||||
if not self._i18n_filter_added:
|
||||
self.template_env.filters['t'] = server_i18n.create_template_filter()
|
||||
self.template_env._i18n_filter_added = True
|
||||
self._i18n_filter_added = True
|
||||
|
||||
template = self.template_env.get_template('statistics.html')
|
||||
rendered = template.render(
|
||||
@@ -549,7 +550,7 @@ class StatsRoutes:
|
||||
'error': str(e)
|
||||
}, status=500)
|
||||
|
||||
def _count_unused_models(self, models: List[Dict], usage_data: Dict) -> int:
|
||||
def _count_unused_models(self, models: List[Dict[str, Any]], usage_data: Dict[str, Any]) -> int:
|
||||
"""Count models that have never been used"""
|
||||
used_hashes = set(usage_data.keys())
|
||||
unused_count = 0
|
||||
@@ -560,7 +561,7 @@ class StatsRoutes:
|
||||
|
||||
return unused_count
|
||||
|
||||
def _get_top_used_models(self, usage_data: Dict, model_map: Dict, limit: int) -> List[Dict]:
|
||||
def _get_top_used_models(self, usage_data: Dict[str, Any], model_map: Dict[str, Any], limit: int) -> List[Dict[str, Any]]:
|
||||
"""Get top used models with their metadata"""
|
||||
sorted_usage = sorted(usage_data.items(), key=lambda x: x[1].get('total', 0), reverse=True)
|
||||
|
||||
@@ -578,7 +579,7 @@ class StatsRoutes:
|
||||
|
||||
return top_models
|
||||
|
||||
def _get_usage_timeline(self, usage_data: Dict, days: int) -> List[Dict]:
|
||||
def _get_usage_timeline(self, usage_data: Dict[str, Any], days: int) -> List[Dict[str, Any]]:
|
||||
"""Get usage timeline for the past N days"""
|
||||
timeline = []
|
||||
today = datetime.now()
|
||||
@@ -614,7 +615,7 @@ class StatsRoutes:
|
||||
|
||||
return list(reversed(timeline)) # Oldest to newest
|
||||
|
||||
def _format_size(self, size_bytes: int) -> str:
|
||||
def _format_size(self, size_bytes: float) -> str:
|
||||
"""Format file size in human readable format"""
|
||||
for unit in ['B', 'KB', 'MB', 'GB', 'TB']:
|
||||
if size_bytes < 1024.0:
|
||||
|
||||
+310
-39
@@ -6,7 +6,7 @@ import shutil
|
||||
import tempfile
|
||||
import asyncio
|
||||
from aiohttp import web, ClientError
|
||||
from typing import Dict, List
|
||||
from typing import Any, Dict, List, cast
|
||||
|
||||
from ..utils.settings_paths import ensure_settings_file
|
||||
from ..services.downloader import get_downloader
|
||||
@@ -38,6 +38,84 @@ def _clean_excludes() -> List[str]:
|
||||
return excludes
|
||||
|
||||
|
||||
def _stage_preserved_items(plugin_root: str) -> tuple[str, list[str]]:
|
||||
"""Move preserved user-data items to a temp directory outside *plugin_root*.
|
||||
|
||||
This ensures that ``git reset --hard``, ``git clean -fd``, and ZIP-based
|
||||
replacement cannot touch these files even when ``-e`` exclusion patterns
|
||||
are mishandled (e.g. on Windows where forward-slash patterns may not
|
||||
match backslash-prefixed paths in some Git builds, or where file locks
|
||||
prevent deletion/recreation).
|
||||
|
||||
Returns:
|
||||
``(backup_root, staged_names)``: the temp directory path and the
|
||||
list of item names that were successfully moved.
|
||||
"""
|
||||
backup_root = tempfile.mkdtemp(prefix='lora_manager_update_')
|
||||
staged: list[str] = []
|
||||
for name in _PRESERVE_DIRS:
|
||||
src = os.path.join(plugin_root, name)
|
||||
if not os.path.lexists(src):
|
||||
continue
|
||||
dst = os.path.join(backup_root, name)
|
||||
try:
|
||||
shutil.move(src, dst)
|
||||
staged.append(name)
|
||||
logger.debug("Staged '%s' for update safety", name)
|
||||
except OSError:
|
||||
# ``shutil.move`` may fail on Windows if a file handle inside
|
||||
# the directory is still open (e.g. a SQLite WAL file). Fall
|
||||
# back to copy-then-remove.
|
||||
logger.debug("Move failed for '%s', falling back to copy", name)
|
||||
try:
|
||||
if os.path.isdir(src) and not os.path.islink(src):
|
||||
shutil.copytree(src, dst, symlinks=True)
|
||||
shutil.rmtree(src, ignore_errors=True)
|
||||
else:
|
||||
shutil.copy2(src, dst)
|
||||
os.remove(src)
|
||||
staged.append(name)
|
||||
logger.info("Copied (then removed) '%s' for update safety", name)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Could not stage '%s': %s (will rely on git -e / skip lists)", name, exc
|
||||
)
|
||||
return backup_root, staged
|
||||
|
||||
|
||||
def _restore_preserved_items(plugin_root: str, backup_root: str, staged: list[str]) -> None:
|
||||
"""Move staged items back from *backup_root* into *plugin_root*.
|
||||
|
||||
Any leftover placeholder at the destination (created by git checkout or
|
||||
ZIP extraction) is removed before the move.
|
||||
"""
|
||||
for name in staged:
|
||||
src = os.path.join(backup_root, name)
|
||||
dst = os.path.join(plugin_root, name)
|
||||
try:
|
||||
if os.path.lexists(dst):
|
||||
if os.path.isdir(dst) and not os.path.islink(dst):
|
||||
shutil.rmtree(dst, ignore_errors=True)
|
||||
else:
|
||||
os.remove(dst)
|
||||
shutil.move(src, dst)
|
||||
logger.debug("Restored '%s' after update", name)
|
||||
except OSError:
|
||||
logger.debug("Move failed restoring '%s', falling back to copy", name)
|
||||
try:
|
||||
if os.path.isdir(src) and not os.path.islink(src):
|
||||
shutil.copytree(src, dst, symlinks=True, dirs_exist_ok=True)
|
||||
shutil.rmtree(src, ignore_errors=True)
|
||||
else:
|
||||
shutil.copy2(src, dst)
|
||||
os.remove(src)
|
||||
logger.info("Copied '%s' back after update", name)
|
||||
except Exception as exc:
|
||||
logger.error("Failed to restore '%s': %s", name, exc)
|
||||
shutil.rmtree(backup_root, ignore_errors=True)
|
||||
|
||||
|
||||
|
||||
class UpdateRoutes:
|
||||
"""Routes for handling plugin update checks"""
|
||||
|
||||
@@ -47,6 +125,7 @@ class UpdateRoutes:
|
||||
app.router.add_get('/api/lm/check-updates', UpdateRoutes.check_updates)
|
||||
app.router.add_get('/api/lm/version-info', UpdateRoutes.get_version_info)
|
||||
app.router.add_post('/api/lm/perform-update', UpdateRoutes.perform_update)
|
||||
app.router.add_post('/api/lm/switch-channel', UpdateRoutes.switch_channel)
|
||||
|
||||
@staticmethod
|
||||
async def check_updates(request):
|
||||
@@ -65,10 +144,17 @@ class UpdateRoutes:
|
||||
|
||||
# Fetch remote version from GitHub
|
||||
if nightly:
|
||||
remote_version, changelog = await UpdateRoutes._get_nightly_version()
|
||||
releases = None
|
||||
local_hash = git_info.get('short_hash', '')
|
||||
nightly_version, releases_result = await asyncio.gather(
|
||||
UpdateRoutes._get_nightly_version(local_hash),
|
||||
UpdateRoutes._get_remote_version()
|
||||
)
|
||||
remote_version, _, behind_by, commit_date = nightly_version
|
||||
_, changelog, releases = releases_result
|
||||
else:
|
||||
remote_version, changelog, releases = await UpdateRoutes._get_remote_version()
|
||||
behind_by = 0
|
||||
commit_date = ''
|
||||
|
||||
# Compare versions
|
||||
if nightly:
|
||||
@@ -81,6 +167,10 @@ class UpdateRoutes:
|
||||
remote_version.replace('v', '')
|
||||
)
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
plugin_root = os.path.dirname(os.path.dirname(current_dir))
|
||||
has_git = os.path.exists(os.path.join(plugin_root, '.git'))
|
||||
|
||||
response_data = {
|
||||
'success': True,
|
||||
'current_version': local_version,
|
||||
@@ -88,13 +178,13 @@ class UpdateRoutes:
|
||||
'update_available': update_available,
|
||||
'changelog': changelog,
|
||||
'git_info': git_info,
|
||||
'nightly': nightly
|
||||
'nightly': nightly,
|
||||
'has_git': has_git,
|
||||
'releases': releases,
|
||||
'behind_by': behind_by,
|
||||
'commit_date': commit_date
|
||||
}
|
||||
|
||||
# Include releases list for stable mode
|
||||
if releases is not None:
|
||||
response_data['releases'] = releases
|
||||
|
||||
return web.json_response(response_data)
|
||||
|
||||
except NETWORK_EXCEPTIONS as e:
|
||||
@@ -126,9 +216,14 @@ class UpdateRoutes:
|
||||
# Format: version-short_hash
|
||||
version_string = f"{local_version}-{short_hash}"
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
plugin_root = os.path.dirname(os.path.dirname(current_dir))
|
||||
has_git = os.path.exists(os.path.join(plugin_root, '.git'))
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'version': version_string
|
||||
'version': version_string,
|
||||
'has_git': has_git
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
@@ -156,20 +251,22 @@ class UpdateRoutes:
|
||||
if os.path.exists(settings_path):
|
||||
with open(settings_path, 'r', encoding='utf-8') as f:
|
||||
settings_backup = f.read()
|
||||
logger.info("Backed up settings.json")
|
||||
logger.debug("Backed up settings.json (%d bytes)", len(settings_backup))
|
||||
|
||||
staged_backup_dir, staged_items = _stage_preserved_items(plugin_root)
|
||||
try:
|
||||
git_folder = os.path.join(plugin_root, '.git')
|
||||
if os.path.exists(git_folder):
|
||||
# Git update
|
||||
success, new_version = await UpdateRoutes._perform_git_update(plugin_root, nightly)
|
||||
else:
|
||||
# Fallback: Download ZIP and replace files
|
||||
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
||||
finally:
|
||||
_restore_preserved_items(plugin_root, staged_backup_dir, staged_items)
|
||||
|
||||
if settings_backup and success:
|
||||
with open(settings_path, 'w', encoding='utf-8') as f:
|
||||
f.write(settings_backup)
|
||||
logger.info("Restored settings.json")
|
||||
logger.debug("Restored settings.json content (%d bytes)", len(settings_backup))
|
||||
|
||||
if success:
|
||||
return web.json_response({
|
||||
@@ -190,6 +287,164 @@ class UpdateRoutes:
|
||||
'error': str(e)
|
||||
})
|
||||
|
||||
@staticmethod
|
||||
async def switch_channel(request):
|
||||
"""
|
||||
Switch between release and nightly update channels.
|
||||
|
||||
ZIP/CNR install → Nightly: git init + checkout main (one-way upgrade)
|
||||
Git install → Release: git checkout latest tag (.git preserved)
|
||||
ZIP/CNR install → Release: ZIP download (no .git, stays in ZIP mode)
|
||||
Git install → Nightly: git checkout main + pull
|
||||
"""
|
||||
try:
|
||||
body = await request.json() if request.has_body else {}
|
||||
channel = body.get('channel', '')
|
||||
|
||||
if channel not in ('release', 'nightly'):
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': f'Invalid channel: {channel}. Must be "release" or "nightly".'
|
||||
})
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
plugin_root = os.path.dirname(os.path.dirname(current_dir))
|
||||
|
||||
settings_path = ensure_settings_file(logger)
|
||||
settings_backup = None
|
||||
if os.path.exists(settings_path):
|
||||
with open(settings_path, 'r', encoding='utf-8') as f:
|
||||
settings_backup = f.read()
|
||||
logger.debug("Backed up settings.json before channel switch (%d bytes)", len(settings_backup))
|
||||
|
||||
staged_backup_dir, staged_items = _stage_preserved_items(plugin_root)
|
||||
try:
|
||||
git_folder = os.path.join(plugin_root, '.git')
|
||||
|
||||
if channel == 'nightly':
|
||||
git_backup = None
|
||||
if os.path.exists(git_folder):
|
||||
git_backup = UpdateRoutes._backup_git(git_folder, 'nightly')
|
||||
|
||||
success = False
|
||||
new_version = ''
|
||||
try:
|
||||
if os.path.exists(git_folder):
|
||||
success, new_version = await UpdateRoutes._perform_git_update(
|
||||
plugin_root, nightly=True
|
||||
)
|
||||
else:
|
||||
success, new_version = UpdateRoutes._init_git_repo(plugin_root)
|
||||
finally:
|
||||
UpdateRoutes._restore_git(git_backup, git_folder, success, 'nightly')
|
||||
else:
|
||||
success = False
|
||||
new_version = ''
|
||||
if os.path.exists(git_folder):
|
||||
success, new_version = await UpdateRoutes._perform_git_update(
|
||||
plugin_root, nightly=False
|
||||
)
|
||||
else:
|
||||
tracking_file = os.path.join(plugin_root, '.tracking')
|
||||
if os.path.exists(tracking_file):
|
||||
os.remove(tracking_file)
|
||||
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
||||
finally:
|
||||
_restore_preserved_items(plugin_root, staged_backup_dir, staged_items)
|
||||
|
||||
if settings_backup and success:
|
||||
with open(settings_path, 'w', encoding='utf-8') as f:
|
||||
f.write(settings_backup)
|
||||
logger.debug("Restored settings.json content after channel switch (%d bytes)", len(settings_backup))
|
||||
|
||||
if success:
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'channel': channel,
|
||||
'new_version': new_version,
|
||||
'message': f'Switched to {channel} channel'
|
||||
})
|
||||
else:
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': f'Failed to switch to {channel} channel'
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Failed to switch channel: %s", e, exc_info=True)
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
})
|
||||
|
||||
@staticmethod
|
||||
def _init_git_repo(plugin_root: str) -> tuple[bool, str]:
|
||||
"""
|
||||
Initialize a Git repository in a ZIP-installed plugin folder.
|
||||
Clones the remote history and checks out main branch.
|
||||
"""
|
||||
try:
|
||||
import git
|
||||
except ImportError:
|
||||
logger.error(
|
||||
"GitPython is not available: cannot initialize git repo. "
|
||||
"Install git or set $GIT_PYTHON_GIT_EXECUTABLE to the git binary path."
|
||||
)
|
||||
return False, ""
|
||||
|
||||
clean_excludes = _clean_excludes()
|
||||
|
||||
try:
|
||||
repo = git.Repo.init(plugin_root)
|
||||
origin = repo.create_remote(
|
||||
'origin',
|
||||
'https://github.com/willmiao/ComfyUI-Lora-Manager.git'
|
||||
)
|
||||
origin.fetch()
|
||||
|
||||
repo.create_head('main', origin.refs.main)
|
||||
repo.git.checkout('main', '--force')
|
||||
repo.git.reset('--hard')
|
||||
repo.git.clean('-fd', *clean_excludes)
|
||||
|
||||
tracking_file = os.path.join(plugin_root, '.tracking')
|
||||
if os.path.exists(tracking_file):
|
||||
os.remove(tracking_file)
|
||||
logger.info("Removed .tracking file (now in git mode)")
|
||||
|
||||
new_version = f"main-{repo.head.commit.hexsha[:7]}"
|
||||
logger.info("Initialized git repo on main branch: %s", new_version)
|
||||
return True, new_version
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Failed to initialize git repo: %s", e, exc_info=True)
|
||||
return False, ""
|
||||
|
||||
@staticmethod
|
||||
def _backup_git(git_folder, label):
|
||||
try:
|
||||
backup_dir = tempfile.mkdtemp()
|
||||
backup = os.path.join(backup_dir, '.git')
|
||||
shutil.copytree(git_folder, backup)
|
||||
logger.info("Backed up .git before switching to %s", label)
|
||||
return backup
|
||||
except Exception as e:
|
||||
logger.error("Failed to backup .git before %s switch: %s", label, e)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _restore_git(git_backup, git_folder, success, label):
|
||||
if git_backup and not success:
|
||||
try:
|
||||
if os.path.exists(git_folder):
|
||||
shutil.rmtree(git_folder)
|
||||
shutil.copytree(git_backup, git_folder)
|
||||
logger.info("Restored .git after failed %s switch", label)
|
||||
except Exception as e:
|
||||
logger.error("Failed to restore .git after %s switch: %s", label, e)
|
||||
if git_backup:
|
||||
shutil.rmtree(os.path.dirname(git_backup), ignore_errors=True)
|
||||
|
||||
@staticmethod
|
||||
async def _download_and_replace_zip(plugin_root: str) -> tuple[bool, str]:
|
||||
"""
|
||||
@@ -213,8 +468,9 @@ class UpdateRoutes:
|
||||
logger.error(f"Failed to fetch release info: {data}")
|
||||
return False, ""
|
||||
|
||||
zip_url = data.get("zipball_url")
|
||||
version = data.get("tag_name", "unknown")
|
||||
release_payload = cast(dict[str, Any], data)
|
||||
zip_url = release_payload.get("zipball_url", "")
|
||||
version = release_payload.get("tag_name", "unknown")
|
||||
|
||||
# Download ZIP to temporary file
|
||||
with tempfile.NamedTemporaryFile(delete=False, suffix=".zip") as tmp_zip:
|
||||
@@ -244,8 +500,7 @@ class UpdateRoutes:
|
||||
except Exception:
|
||||
logger.debug("Could not close downloaded-version history database", exc_info=True)
|
||||
|
||||
# Skip settings.json, civitai, model cache and runtime cache folders
|
||||
UpdateRoutes._clean_plugin_folder(plugin_root, skip_files=['settings.json', 'civitai', 'model_cache', 'cache', 'wildcards', 'backups', 'stats'])
|
||||
UpdateRoutes._clean_plugin_folder(plugin_root, skip_files=list(_PRESERVE_DIRS))
|
||||
|
||||
# Extract ZIP to temp dir
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
@@ -255,7 +510,7 @@ class UpdateRoutes:
|
||||
extracted_root = next(os.scandir(tmp_dir)).path
|
||||
|
||||
# Copy files, skipping user data that should be preserved
|
||||
skip_items = {'settings.json', 'civitai', 'wildcards', 'backups', 'stats'}
|
||||
skip_items = set(_PRESERVE_DIRS)
|
||||
for item in os.listdir(extracted_root):
|
||||
if item in skip_items:
|
||||
continue
|
||||
@@ -272,7 +527,7 @@ class UpdateRoutes:
|
||||
# for ComfyUI Manager to work properly
|
||||
tracking_info_file = os.path.join(plugin_root, '.tracking')
|
||||
tracking_files = []
|
||||
skip_tracked = {'civitai', 'wildcards', 'backups', 'stats'}
|
||||
skip_tracked = set(_PRESERVE_DIRS) - {'settings.json'}
|
||||
for root, dirs, files in os.walk(extracted_root):
|
||||
# Skip user data directories and their contents
|
||||
rel_root = os.path.relpath(root, extracted_root)
|
||||
@@ -296,6 +551,7 @@ class UpdateRoutes:
|
||||
logger.error(f"ZIP update failed: {e}", exc_info=True)
|
||||
return False, ""
|
||||
|
||||
@staticmethod
|
||||
def _clean_plugin_folder(plugin_root, skip_files=None):
|
||||
skip_files = skip_files or []
|
||||
for item in os.listdir(plugin_root):
|
||||
@@ -308,41 +564,56 @@ class UpdateRoutes:
|
||||
os.remove(path)
|
||||
|
||||
@staticmethod
|
||||
async def _get_nightly_version() -> tuple[str, List[str]]:
|
||||
"""
|
||||
Fetch latest commit from main branch
|
||||
"""
|
||||
async def _get_nightly_version(local_hash: str = "") -> tuple[str, List[str], int, str]:
|
||||
repo_owner = "willmiao"
|
||||
repo_name = "ComfyUI-Lora-Manager"
|
||||
|
||||
# Use GitHub API to fetch the latest commit from main branch
|
||||
github_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/commits/main"
|
||||
|
||||
try:
|
||||
downloader = await get_downloader()
|
||||
success, data = await downloader.make_request('GET', github_url, custom_headers={'Accept': 'application/vnd.github+json'})
|
||||
success, data = await downloader.make_request(
|
||||
'GET', github_url,
|
||||
custom_headers={'Accept': 'application/vnd.github+json'}
|
||||
)
|
||||
|
||||
if not success:
|
||||
logger.warning(f"Failed to fetch GitHub commit: {data}")
|
||||
return "main", []
|
||||
logger.warning("Failed to fetch GitHub commit: %s", data)
|
||||
return "main", [], 0, ""
|
||||
|
||||
commit_sha = data.get('sha', '')[:7] # Short hash
|
||||
commit_message = data.get('commit', {}).get('message', '')
|
||||
commit_payload = cast(dict[str, Any], data)
|
||||
commit_sha = commit_payload.get('sha', '')[:7]
|
||||
commit_message = commit_payload.get('commit', {}).get('message', '')
|
||||
commit_date = commit_payload.get('commit', {}).get('committer', {}).get('date', '')[:10]
|
||||
|
||||
# Format as "main-{short_hash}"
|
||||
version = f"main-{commit_sha}"
|
||||
|
||||
# Use commit message as changelog
|
||||
changelog = [commit_message] if commit_message else []
|
||||
|
||||
return version, changelog
|
||||
behind_by = 0
|
||||
if local_hash and local_hash not in ('unknown', 'stable'):
|
||||
compare_url = (
|
||||
f"https://api.github.com/repos/{repo_owner}/{repo_name}"
|
||||
f"/compare/{local_hash}...main"
|
||||
)
|
||||
c_ok, c_data = await downloader.make_request(
|
||||
'GET', compare_url,
|
||||
custom_headers={'Accept': 'application/vnd.github+json'}
|
||||
)
|
||||
if c_ok:
|
||||
compare_payload = cast(dict[str, Any], c_data)
|
||||
if compare_payload.get('status') in ('ahead', 'diverged'):
|
||||
behind_by = compare_payload.get('ahead_by', 0)
|
||||
else:
|
||||
behind_by = compare_payload.get('behind_by', 0)
|
||||
|
||||
return version, changelog, behind_by, commit_date
|
||||
|
||||
except NETWORK_EXCEPTIONS as e:
|
||||
logger.warning("Unable to reach GitHub for nightly version: %s", e)
|
||||
return "main", []
|
||||
return "main", [], 0, ""
|
||||
except Exception as e:
|
||||
logger.error(f"Error fetching nightly version: {e}", exc_info=True)
|
||||
return "main", []
|
||||
logger.error("Error fetching nightly version: %s", e, exc_info=True)
|
||||
return "main", [], 0, ""
|
||||
|
||||
@staticmethod
|
||||
def _compare_nightly_versions(local_git_info: Dict[str, str], remote_version: str) -> bool:
|
||||
@@ -438,7 +709,7 @@ class UpdateRoutes:
|
||||
logger.info(f"Successfully updated to {new_version}")
|
||||
return True, new_version
|
||||
|
||||
except git.exc.GitError as e:
|
||||
except git.exc.GitError as e: # pyright: ignore[reportAttributeAccessIssue]
|
||||
logger.error(f"Git error during update: {e}")
|
||||
return False, ""
|
||||
except Exception as e:
|
||||
@@ -499,7 +770,7 @@ class UpdateRoutes:
|
||||
return git_info
|
||||
|
||||
@staticmethod
|
||||
async def _get_remote_version() -> tuple[str, List[str], List[Dict]]:
|
||||
async def _get_remote_version() -> tuple[str, List[str], List[Dict[str, Any]]]:
|
||||
"""
|
||||
Fetch remote version from GitHub
|
||||
Returns:
|
||||
@@ -521,7 +792,7 @@ class UpdateRoutes:
|
||||
|
||||
# Parse releases
|
||||
releases = []
|
||||
for i, release in enumerate(data):
|
||||
for i, release in enumerate(cast(list[dict[str, Any]], data)):
|
||||
version = release.get('tag_name', '')
|
||||
if not version.startswith('v'):
|
||||
version = f"v{version}"
|
||||
|
||||
@@ -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
|
||||
@@ -117,7 +117,7 @@ def _render_prompt(template: str, variables: Dict[str, Any]) -> str:
|
||||
Uses simple regex substitution — no Jinja2 dependency needed.
|
||||
"""
|
||||
|
||||
def replace(match: re.Match) -> str:
|
||||
def replace(match: re.Match[str]) -> str:
|
||||
key = match.group(1).strip()
|
||||
value = variables.get(key, "")
|
||||
if isinstance(value, (dict, list)):
|
||||
|
||||
@@ -295,7 +295,7 @@ class PostProcessor:
|
||||
normalises every tag to lowercase for case-insensitive dedup.
|
||||
"""
|
||||
merged: List[str] = []
|
||||
seen: set = set()
|
||||
seen: set[str] = set()
|
||||
for tag in list(existing) + list(new):
|
||||
t = tag.strip().lower()
|
||||
if t and t not in seen:
|
||||
|
||||
@@ -49,7 +49,7 @@ _FRONTMATTER_RE = re.compile(
|
||||
)
|
||||
|
||||
|
||||
def _parse_skill_file(path: Path) -> tuple[dict, str]:
|
||||
def _parse_skill_file(path: Path) -> tuple[dict[str, Any], str]:
|
||||
"""Read a prompt definition file (``prompt.md`` or legacy ``SKILL.md``) and
|
||||
return (frontmatter_dict, body_text).
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ from __future__ import annotations
|
||||
|
||||
import html as html_module
|
||||
import re
|
||||
from typing import List, Tuple
|
||||
from typing import Any, List, Tuple
|
||||
|
||||
|
||||
_REPO_URL_PATTERN = re.compile(r"https?://huggingface\.co/([^/]+/[^/]+)")
|
||||
@@ -18,10 +18,10 @@ _REPO_URL_PATTERN = re.compile(r"https?://huggingface\.co/([^/]+/[^/]+)")
|
||||
def extract_simple_markdown_images(
|
||||
markdown_text: str,
|
||||
repo: str,
|
||||
existing_urls: set | None = None,
|
||||
existing_urls: set[str] | None = None,
|
||||
default_width: int = 512,
|
||||
default_height: int = 512,
|
||||
) -> list[dict]:
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Extract standalone markdown images from the README body.
|
||||
|
||||
Matches ```` on lines that are NOT part of a markdown table
|
||||
@@ -36,8 +36,8 @@ def extract_simple_markdown_images(
|
||||
return []
|
||||
|
||||
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
||||
images: list[dict] = []
|
||||
seen_urls: set = set(existing_urls) if existing_urls else set()
|
||||
images: list[dict[str, Any]] = []
|
||||
seen_urls: set[str] = set(existing_urls) if existing_urls else set()
|
||||
|
||||
# Collect lines that are NOT inside fenced code blocks
|
||||
lines = markdown_text.split("\n")
|
||||
@@ -86,10 +86,10 @@ def extract_simple_markdown_images(
|
||||
def extract_html_img_tags(
|
||||
markdown_text: str,
|
||||
repo: str,
|
||||
existing_urls: set | None = None,
|
||||
existing_urls: set[str] | None = None,
|
||||
default_width: int = 512,
|
||||
default_height: int = 512,
|
||||
) -> list[dict]:
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Extract image URLs from HTML ``<img src=\"...\">`` tags in the README.
|
||||
|
||||
Many HF collection repos (e.g. ``deadman44/Z-Image_LoRA``) use raw HTML
|
||||
@@ -103,8 +103,8 @@ def extract_html_img_tags(
|
||||
return []
|
||||
|
||||
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
||||
images: list[dict] = []
|
||||
seen_urls: set = set(existing_urls) if existing_urls else set()
|
||||
images: list[dict[str, Any]] = []
|
||||
seen_urls: set[str] = set(existing_urls) if existing_urls else set()
|
||||
|
||||
for m in re.finditer(
|
||||
r'<img\s[^>]*src=\"([^\"]+)\"',
|
||||
@@ -175,7 +175,7 @@ def extract_gallery_images(
|
||||
repo: str,
|
||||
default_width: int = 512,
|
||||
default_height: int = 512,
|
||||
) -> List[dict]:
|
||||
) -> List[dict[str, Any]]:
|
||||
"""Extract widget/gallery images from the YAML frontmatter of a HF README.
|
||||
|
||||
Args:
|
||||
@@ -196,7 +196,7 @@ def extract_gallery_images(
|
||||
if not frontmatter:
|
||||
return []
|
||||
|
||||
images: List[dict] = []
|
||||
images: List[dict[str, Any]] = []
|
||||
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
||||
w = default_width or 512
|
||||
h = default_height or 512
|
||||
@@ -258,7 +258,7 @@ def extract_gallery_images(
|
||||
text = raw_text
|
||||
|
||||
if url:
|
||||
image: dict = {
|
||||
image: dict[str, Any] = {
|
||||
"url": url,
|
||||
"type": "image",
|
||||
"nsfwLevel": 0,
|
||||
@@ -276,10 +276,10 @@ def extract_gallery_images(
|
||||
def extract_gallery_table_images(
|
||||
markdown_text: str,
|
||||
repo: str,
|
||||
existing_urls: set | None = None,
|
||||
existing_urls: set[str] | None = None,
|
||||
default_width: int = 512,
|
||||
default_height: int = 512,
|
||||
) -> list[dict]:
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Extract images from ``| Preview | Prompt |`` markdown gallery tables.
|
||||
|
||||
Many HF READMEs include a sample-gallery table in the body (outside
|
||||
@@ -295,8 +295,8 @@ def extract_gallery_table_images(
|
||||
return []
|
||||
|
||||
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
||||
images: list[dict] = []
|
||||
seen_urls: set = set(existing_urls) if existing_urls else set()
|
||||
images: list[dict[str, Any]] = []
|
||||
seen_urls: set[str] = set(existing_urls) if existing_urls else set()
|
||||
lines = markdown_text.split("\n")
|
||||
n = len(lines)
|
||||
i = 0
|
||||
@@ -514,7 +514,7 @@ def _strip_standalone_images(text: str) -> str:
|
||||
URL was stripped entirely, making it impossible for the LLM to return
|
||||
a ``preview_url`` for repos that use HTML ``<img>`` tags exclusively.
|
||||
"""
|
||||
def _img_to_md(match: re.Match) -> str:
|
||||
def _img_to_md(match: re.Match[str]) -> str:
|
||||
"""Convert an ``<img>`` tag to markdown image syntax ````."""
|
||||
tag = match.group(0)
|
||||
src_m = re.search(r'src="([^"]+)"', tag) or re.search(r"src='([^']+)'", tag)
|
||||
@@ -942,7 +942,7 @@ def _strip_badge_images(text: str) -> str:
|
||||
"twitter", "colab", "gradio", "space",
|
||||
)
|
||||
|
||||
def _should_remove(m: re.Match) -> str:
|
||||
def _should_remove(m: re.Match[str]) -> str:
|
||||
alt = (m.group(1) or "").lower()
|
||||
for kw in badge_keywords:
|
||||
if kw in alt:
|
||||
|
||||
+221
-16
@@ -1,3 +1,7 @@
|
||||
# pyright: reportImportCycles=false
|
||||
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||
# import cycles. Breaking them would require an architectural refactor.
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
@@ -7,6 +11,7 @@ import os
|
||||
import secrets
|
||||
import shutil
|
||||
import socket
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
@@ -20,10 +25,43 @@ from .settings_manager import get_settings_manager
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Maximum times the download poll loop will re-schedule a transfer after it
|
||||
# is lost (daemon restart / RPC outage) before failing the download.
|
||||
MAX_TRANSFER_RECOVERY_ATTEMPTS = 2
|
||||
|
||||
# stderr lines matching these markers indicate a disk write failure inside
|
||||
# aria2 (piece cache flush or raw file write). They are promoted to INFO so
|
||||
# the root cause (disk full, permission denied, file locked by another
|
||||
# process, ...) is visible in the default logs; all other stderr output stays
|
||||
# at DEBUG to avoid noise.
|
||||
_DISK_WRITE_ERROR_MARKERS = (
|
||||
# aria2 wrapper messages (write disk cache flush path)
|
||||
"write disk cache flush failure",
|
||||
"error when trying to flush write cache",
|
||||
"failed to write into the file",
|
||||
"failed to open the file",
|
||||
"failed to seek the file",
|
||||
# underlying root-cause phrases reported via "cause: ..." (POSIX + Windows)
|
||||
"no space left on device",
|
||||
"not enough space on the disk",
|
||||
"input/output error",
|
||||
"permission denied",
|
||||
"access is denied",
|
||||
"disk quota exceeded",
|
||||
"used by another process",
|
||||
"sharing violation",
|
||||
)
|
||||
|
||||
# Minimum interval between INFO-level reports of the same stderr line so a
|
||||
# repeated failure (e.g. aria2 retrying against a full disk) does not spam
|
||||
# the log.
|
||||
STDERR_ERROR_REPORT_INTERVAL = 60.0
|
||||
|
||||
|
||||
def _try_certifi_ca_path() -> str | None:
|
||||
"""Return the certifi CA bundle path if available, else None."""
|
||||
try:
|
||||
import certifi # type: ignore[import-untyped]
|
||||
import certifi # pyright: ignore[reportMissingTypeStubs]
|
||||
|
||||
path = certifi.where()
|
||||
if os.path.isfile(path):
|
||||
@@ -44,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):
|
||||
"""Raised when aria2 integration fails."""
|
||||
|
||||
@@ -81,10 +130,12 @@ class Aria2Downloader:
|
||||
self._rpc_session: Optional[aiohttp.ClientSession] = None
|
||||
self._rpc_session_lock = asyncio.Lock()
|
||||
self._process_lock = asyncio.Lock()
|
||||
self._register_lock = asyncio.Lock()
|
||||
self._transfers: Dict[str, Aria2Transfer] = {}
|
||||
self._poll_interval = 0.5
|
||||
self._state_store = Aria2TransferStateStore()
|
||||
self._stderr_reader_task: Optional[asyncio.Task] = None
|
||||
self._stderr_reader_task: Optional[asyncio.Task[Any]] = None
|
||||
self._stderr_error_report: Dict[str, float] = {}
|
||||
|
||||
@property
|
||||
def is_running(self) -> bool:
|
||||
@@ -99,26 +150,61 @@ class Aria2Downloader:
|
||||
progress_callback=None,
|
||||
headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[bool, str]:
|
||||
"""Download a file using aria2 RPC and wait for completion."""
|
||||
"""Download a file using aria2 RPC and wait for completion.
|
||||
|
||||
The poll loop is self-healing: when the in-memory transfer entry
|
||||
disappears (e.g. another download restarted the daemon and
|
||||
``close()`` cleared ``_transfers``) or the RPC becomes unreachable,
|
||||
the transfer is re-scheduled with ``continue=true`` so the download
|
||||
resumes from the on-disk ``.aria2`` control file. The same
|
||||
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``.
|
||||
"""
|
||||
|
||||
await self._ensure_process()
|
||||
save_path = os.path.abspath(save_path)
|
||||
|
||||
async with self._register_lock:
|
||||
transfer = self._transfers.get(download_id)
|
||||
if transfer is None or os.path.abspath(transfer.save_path) != save_path:
|
||||
gid = await self._schedule_download(
|
||||
transfer = await self._register_transfer(
|
||||
url,
|
||||
save_path,
|
||||
download_id=download_id,
|
||||
headers=headers,
|
||||
)
|
||||
transfer = Aria2Transfer(gid=gid, save_path=save_path)
|
||||
self._transfers[download_id] = transfer
|
||||
|
||||
recovery_attempts = 0
|
||||
try:
|
||||
while True:
|
||||
try:
|
||||
status = await self._get_status_with_retry(download_id)
|
||||
except Aria2Error:
|
||||
status = None
|
||||
|
||||
if status is None:
|
||||
if recovery_attempts >= MAX_TRANSFER_RECOVERY_ATTEMPTS:
|
||||
return False, "aria2 download not found"
|
||||
recovery_attempts += 1
|
||||
logger.warning(
|
||||
"aria2 transfer %s lost; re-scheduling with resume "
|
||||
"(attempt %d/%d)",
|
||||
download_id,
|
||||
recovery_attempts,
|
||||
MAX_TRANSFER_RECOVERY_ATTEMPTS,
|
||||
)
|
||||
await asyncio.sleep(1.0)
|
||||
await self._ensure_process()
|
||||
async with self._register_lock:
|
||||
transfer = await self._register_transfer(
|
||||
url,
|
||||
save_path,
|
||||
download_id=download_id,
|
||||
headers=headers,
|
||||
)
|
||||
continue
|
||||
|
||||
snapshot = self._build_progress_snapshot(status)
|
||||
if progress_callback is not None:
|
||||
@@ -129,12 +215,43 @@ class Aria2Downloader:
|
||||
completed_path = self._resolve_completed_path(status, save_path)
|
||||
return True, completed_path
|
||||
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":
|
||||
return False, "Download was cancelled"
|
||||
|
||||
await asyncio.sleep(self._poll_interval)
|
||||
finally:
|
||||
current = self._transfers.get(download_id)
|
||||
if current is not None and current.gid == transfer.gid:
|
||||
self._transfers.pop(download_id, None)
|
||||
|
||||
async def _get_status_with_retry(
|
||||
@@ -143,8 +260,9 @@ class Aria2Downloader:
|
||||
"""Call get_status with retry for transient RPC failures.
|
||||
|
||||
Only retries on :exc:`Aria2Error` (RPC-level failure). Returns
|
||||
``None`` immediately when the download_id is not tracked (a missing
|
||||
transfer is not a transient condition, so retrying is pointless).
|
||||
``None`` immediately when the transfer is not tracked or its GID is
|
||||
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,
|
||||
because aria2 may be temporarily busy (e.g. finalizing multiple
|
||||
@@ -190,7 +308,7 @@ class Aria2Downloader:
|
||||
download_id,
|
||||
)
|
||||
|
||||
options: Dict[str, str] = {
|
||||
options: Dict[str, Any] = {
|
||||
"dir": save_dir,
|
||||
"out": out_name,
|
||||
"continue": "true",
|
||||
@@ -238,8 +356,33 @@ class Aria2Downloader:
|
||||
)
|
||||
return gid
|
||||
|
||||
async def _register_transfer(
|
||||
self,
|
||||
url: str,
|
||||
save_path: str,
|
||||
*,
|
||||
download_id: str,
|
||||
headers: Optional[Dict[str, str]] = None,
|
||||
) -> Aria2Transfer:
|
||||
"""Schedule a download and track it in the in-memory transfer registry."""
|
||||
gid = await self._schedule_download(
|
||||
url,
|
||||
save_path,
|
||||
download_id=download_id,
|
||||
headers=headers,
|
||||
)
|
||||
transfer = Aria2Transfer(gid=gid, save_path=os.path.abspath(save_path))
|
||||
self._transfers[download_id] = transfer
|
||||
return transfer
|
||||
|
||||
async def get_status(self, download_id: str) -> Optional[Dict[str, Any]]:
|
||||
"""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)
|
||||
if transfer is None:
|
||||
@@ -255,8 +398,17 @@ class Aria2Downloader:
|
||||
"files",
|
||||
]
|
||||
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:
|
||||
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
|
||||
|
||||
if isinstance(status, dict):
|
||||
@@ -274,7 +426,9 @@ class Aria2Downloader:
|
||||
"files",
|
||||
]
|
||||
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:
|
||||
message = str(exc)
|
||||
if "cannot be found" in message.lower() or "not found" in message.lower():
|
||||
@@ -341,8 +495,19 @@ class Aria2Downloader:
|
||||
try:
|
||||
await self._rpc_call("aria2.forceRemove", [transfer.gid])
|
||||
except Exception as 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)
|
||||
return {"success": True, "message": "Download cancelled successfully"}
|
||||
|
||||
@@ -385,16 +550,51 @@ class Aria2Downloader:
|
||||
blocks, which freezes the entire ``aria2c`` process — including its
|
||||
RPC handler. This background task reads lines from stderr as they
|
||||
arrive and forwards them to Python's logger.
|
||||
|
||||
Lines that indicate a disk write failure (e.g. the "cause: No space
|
||||
left on device" line that follows "Write disk cache flush failure")
|
||||
are promoted to INFO so the root cause is visible without enabling
|
||||
debug logging; every other line stays at DEBUG to avoid noise.
|
||||
"""
|
||||
try:
|
||||
assert self._process is not None and self._process.stderr is not None
|
||||
async for line in self._process.stderr:
|
||||
text = line.decode("utf-8", errors="replace").rstrip()
|
||||
if text:
|
||||
if self._is_disk_write_error(text):
|
||||
self._report_stderr_error(text)
|
||||
else:
|
||||
logger.debug("aria2 stderr: %s", text)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def _is_disk_write_error(text: str) -> bool:
|
||||
lowered = text.lower()
|
||||
return any(marker in lowered for marker in _DISK_WRITE_ERROR_MARKERS)
|
||||
|
||||
def _report_stderr_error(self, text: str) -> None:
|
||||
"""INFO-log a disk write failure line, rate-limited per line text.
|
||||
|
||||
aria2 re-emits the same error chain on every poll/retry while the
|
||||
underlying condition persists; only the first occurrence within
|
||||
``STDERR_ERROR_REPORT_INTERVAL`` seconds is promoted to INFO.
|
||||
"""
|
||||
now = time.monotonic()
|
||||
last = self._stderr_error_report.get(text)
|
||||
if last is not None and now - last < STDERR_ERROR_REPORT_INTERVAL:
|
||||
logger.debug("aria2 stderr (repeated disk write error): %s", text)
|
||||
return
|
||||
# Drop entries older than the window so the map stays bounded even
|
||||
# during a long disk-full episode (piece indexes change per line).
|
||||
self._stderr_error_report = {
|
||||
line: timestamp
|
||||
for line, timestamp in self._stderr_error_report.items()
|
||||
if now - timestamp < STDERR_ERROR_REPORT_INTERVAL
|
||||
}
|
||||
self._stderr_error_report[text] = now
|
||||
logger.info("aria2 disk write failure: %s", text)
|
||||
|
||||
async def _dispatch_progress(self, callback, snapshot: DownloadProgress) -> None:
|
||||
try:
|
||||
result = callback(snapshot, snapshot)
|
||||
@@ -597,7 +797,9 @@ class Aria2Downloader:
|
||||
|
||||
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:
|
||||
raise Aria2Error("aria2 RPC endpoint is not initialized")
|
||||
|
||||
@@ -628,7 +830,10 @@ class Aria2Downloader:
|
||||
error = body["error"] or {}
|
||||
code = error.get("code") if isinstance(error, dict) else None
|
||||
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",
|
||||
method,
|
||||
response.status,
|
||||
@@ -643,7 +848,7 @@ class Aria2Downloader:
|
||||
raise Aria2Error(status_message or "Unknown aria2 RPC error")
|
||||
|
||||
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",
|
||||
method,
|
||||
response.status,
|
||||
|
||||
@@ -8,7 +8,7 @@ from filename, base_model, and CivitAI version name — no manual tagging requir
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from typing import Dict, List, Set
|
||||
from typing import Any, Dict, List, Set
|
||||
|
||||
# ── Tag category definitions ──────────────────────────────────────────
|
||||
# Each category maps a display label to a regex pattern.
|
||||
@@ -52,7 +52,7 @@ AUTO_TAG_GROUPS = {
|
||||
DEFAULT_ENABLED_GROUPS = {"mode", "video"}
|
||||
|
||||
|
||||
def _collect_sources(model_data: Dict) -> List[str]:
|
||||
def _collect_sources(model_data: Dict[str, Any]) -> List[str]:
|
||||
"""Collect all text sources from model data for tag matching."""
|
||||
sources: List[str] = []
|
||||
|
||||
@@ -73,7 +73,7 @@ def _collect_sources(model_data: Dict) -> List[str]:
|
||||
return sources
|
||||
|
||||
|
||||
def extract_auto_tags(model_data: Dict) -> List[str]:
|
||||
def extract_auto_tags(model_data: Dict[str, Any]) -> List[str]:
|
||||
"""Extract auto-detected tags from model metadata.
|
||||
|
||||
Uses a two-layer approach:
|
||||
|
||||
@@ -0,0 +1,144 @@
|
||||
# pyright: reportImportCycles=false
|
||||
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||
# import cycles. Breaking them would require an architectural refactor.
|
||||
"""Backfill the AutoV3 checked state for models loaded from a persisted snapshot.
|
||||
|
||||
The SQLite persistent cache predates the AutoV3 feature, so entries hydrated
|
||||
from it have a NULL ``autov3`` column (the "not checked yet" state). This
|
||||
service computes the embedded AutoV3 hash for each such model — once per
|
||||
process — and persists it through the scanner's single write path
|
||||
(:meth:`ModelScanner.update_autov3_for_model`), marking every visited row so a
|
||||
subsequent run finds nothing left to do.
|
||||
|
||||
Three-state contract honored here:
|
||||
|
||||
- ``NULL`` (sqlite) / absent (dict) = not checked yet → backfill computes it
|
||||
- ``''`` (sqlite/dict) / JSON null = checked, no value available → never recompute
|
||||
- 12-char lowercase hex = value → never recompute
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
if TYPE_CHECKING: # pragma: no cover - type-check only; runtime imports are local
|
||||
from .model_scanner import ModelScanner
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _resolve_autov3(file_path: str) -> str:
|
||||
"""Resolve the AutoV3 hash for a model file.
|
||||
|
||||
Prefers the Civitai AutoV3 reported for the file whose SHA256 matches
|
||||
(the authoritative value for recipe matching); falls back to the embedded
|
||||
safetensors header hash. Returns ``''`` when neither is available.
|
||||
"""
|
||||
try:
|
||||
metadata_path = f"{os.path.splitext(file_path)[0]}.metadata.json"
|
||||
if os.path.exists(metadata_path):
|
||||
with open(metadata_path, "r", encoding="utf-8") as handle:
|
||||
payload = json.load(handle)
|
||||
if isinstance(payload, dict):
|
||||
from ..utils.models import autov3_from_civitai_files # local import avoids cycles
|
||||
|
||||
sha256 = (payload.get("sha256") or "").lower()
|
||||
civitai_autov3 = autov3_from_civitai_files(payload.get("civitai"), sha256)
|
||||
if civitai_autov3:
|
||||
return civitai_autov3
|
||||
except Exception:
|
||||
pass
|
||||
from ..utils.file_utils import calculate_autov3 # local import avoids cycles
|
||||
|
||||
return calculate_autov3(file_path) or ""
|
||||
|
||||
|
||||
class Autov3BackfillService:
|
||||
"""Compute and persist AutoV3 hashes for models missing a checked state."""
|
||||
|
||||
_instance: Optional["Autov3BackfillService"] = None
|
||||
_instance_lock = threading.Lock()
|
||||
|
||||
def __init__(self) -> None:
|
||||
# Re-entrancy guard per model type: scanners for different model types
|
||||
# initialize concurrently (lora_manager.py), so a global guard would
|
||||
# silently skip every type but the first to start. Each model type
|
||||
# runs its own backfill; a duplicate trigger for the same type no-ops.
|
||||
self._running_types: set[str] = set()
|
||||
|
||||
@classmethod
|
||||
def get_instance(cls) -> "Autov3BackfillService":
|
||||
"""Return the process-wide singleton instance."""
|
||||
if cls._instance is None:
|
||||
with cls._instance_lock:
|
||||
if cls._instance is None:
|
||||
cls._instance = cls()
|
||||
return cls._instance
|
||||
|
||||
async def backfill(self, scanner: "ModelScanner") -> int:
|
||||
"""Compute AutoV3 for every un-checked model of ``scanner.model_type``.
|
||||
|
||||
Each candidate file is read once via :func:`~py.utils.file_utils.calculate_autov3`
|
||||
(cheap: safetensors header only) and the result is persisted through
|
||||
``scanner.update_autov3_for_model``. Files that no longer exist on
|
||||
disk are skipped — they are intentionally NOT marked, because scanner
|
||||
cleanup removes the stale row later.
|
||||
|
||||
Returns:
|
||||
The number of models successfully updated. Never raises; on any
|
||||
failure a warning is logged and ``0`` is returned. A duplicate
|
||||
trigger for a model type that is already being backfilled returns
|
||||
``0`` immediately; different model types run concurrently.
|
||||
"""
|
||||
model_type = scanner.model_type
|
||||
if model_type in self._running_types:
|
||||
return 0
|
||||
self._running_types.add(model_type)
|
||||
try:
|
||||
# Local imports avoid import cycles at module load time.
|
||||
from .persistent_model_cache import get_persistent_cache
|
||||
from ..utils.file_utils import calculate_autov3
|
||||
|
||||
persistent = getattr(scanner, "_persistent_cache", None) or get_persistent_cache()
|
||||
paths = persistent.get_models_missing_autov3(model_type)
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
count = 0
|
||||
for path in paths:
|
||||
# A file that no longer exists must not be marked; scanner
|
||||
# cleanup removes the stale row later. The existence check and
|
||||
# hash resolution run in the executor so the loop stays
|
||||
# responsive to API requests while the backfill iterates a
|
||||
# large library.
|
||||
if not await loop.run_in_executor(None, os.path.exists, path):
|
||||
continue
|
||||
autov3 = await loop.run_in_executor(None, _resolve_autov3, path)
|
||||
if await scanner.update_autov3_for_model(model_type, path, autov3):
|
||||
count += 1
|
||||
|
||||
if paths:
|
||||
logger.info(
|
||||
"AutoV3 backfill: updated %d/%d models for %s",
|
||||
count,
|
||||
len(paths),
|
||||
model_type,
|
||||
)
|
||||
else:
|
||||
# Steady state after the first run: nothing left to backfill.
|
||||
logger.debug("AutoV3 backfill: nothing to process for %s", model_type)
|
||||
return count
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"AutoV3 backfill failed for %s: %s",
|
||||
getattr(scanner, "model_type", "?"),
|
||||
exc,
|
||||
)
|
||||
return 0
|
||||
finally:
|
||||
self._running_types.discard(model_type)
|
||||
@@ -1,3 +1,7 @@
|
||||
# pyright: reportImportCycles=false
|
||||
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||
# import cycles. Breaking them would require an architectural refactor.
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
from abc import ABC, abstractmethod
|
||||
import asyncio
|
||||
import re
|
||||
from typing import Any, Dict, List, Optional, Type, TYPE_CHECKING
|
||||
import random
|
||||
from typing import Any, Awaitable, Dict, List, Optional, Type, Union, TYPE_CHECKING, cast
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
@@ -69,24 +70,24 @@ class BaseModelService(ABC):
|
||||
page: int,
|
||||
page_size: int,
|
||||
sort_by: str = "name",
|
||||
folder: str = None,
|
||||
folder_include: list = None,
|
||||
folder_exclude: list = None,
|
||||
search: str = None,
|
||||
folder: str | None = None,
|
||||
folder_include: list[str] | None = None,
|
||||
folder_exclude: list[str] | None = None,
|
||||
search: str | None = None,
|
||||
fuzzy_search: bool = False,
|
||||
base_models: list = None,
|
||||
model_types: list = None,
|
||||
base_models: list[str] | None = None,
|
||||
model_types: list[str] | None = None,
|
||||
tags: Optional[Dict[str, str]] = None,
|
||||
auto_tags: Optional[Dict[str, str]] = None,
|
||||
search_options: dict = None,
|
||||
hash_filters: dict = None,
|
||||
search_options: dict[str, Any] | None = None,
|
||||
hash_filters: dict[str, Any] | None = None,
|
||||
favorites_only: bool = False,
|
||||
update_available_only: bool = False,
|
||||
credit_required: Optional[bool] = None,
|
||||
allow_selling_generated_content: Optional[bool] = None,
|
||||
tag_logic: str = "any",
|
||||
**kwargs,
|
||||
) -> Dict:
|
||||
) -> Dict[str, Any]:
|
||||
"""Get paginated and filtered model data"""
|
||||
overall_start = time.perf_counter()
|
||||
|
||||
@@ -109,12 +110,15 @@ class BaseModelService(ABC):
|
||||
if civitai_model_id is not None:
|
||||
sorted_data = [
|
||||
item for item in sorted_data
|
||||
if self._extract_model_id(item) == civitai_model_id
|
||||
if self._extract_group_key(item) == civitai_model_id
|
||||
]
|
||||
# VLM mode: always sort by version ID descending (newest version first),
|
||||
# regardless of the current sort_by preference.
|
||||
# Fall back to modified timestamp for non-CivitAI sources.
|
||||
sorted_data.sort(
|
||||
key=lambda x: self._extract_version_id(x) or 0,
|
||||
key=lambda x: self._extract_version_id(x)
|
||||
or x.get("modified", 0)
|
||||
or 0,
|
||||
reverse=True,
|
||||
)
|
||||
|
||||
@@ -129,18 +133,21 @@ class BaseModelService(ABC):
|
||||
ufs = self.settings.get("version_grouping", "same_base")
|
||||
group_by_base = ufs == "same_base"
|
||||
|
||||
dedup_map = {} # (modelId [,base_model]) -> (item, version_id)
|
||||
dedup_map = {} # (modelId [,base_model]) -> (item, version_or_modified)
|
||||
version_counter = {} # same-key -> count
|
||||
standalone = []
|
||||
for item in sorted_data:
|
||||
mid = self._extract_model_id(item)
|
||||
mid = self._extract_group_key(item)
|
||||
if mid is None:
|
||||
standalone.append(item)
|
||||
continue
|
||||
key = (mid, item.get("base_model") or "") if group_by_base else mid
|
||||
# Count all versions per key
|
||||
version_counter[key] = version_counter.get(key, 0) + 1
|
||||
vid = self._extract_version_id(item) or 0
|
||||
# Prefer CivitAI version_id; fall back to modified timestamp
|
||||
vid = self._extract_version_id(item)
|
||||
if vid is None:
|
||||
vid = item.get("modified", 0) or 0
|
||||
if key not in dedup_map or vid > dedup_map[key][1]:
|
||||
dedup_map[key] = (item, vid)
|
||||
# Attach version_count to each surviving grouped item (shallow copy
|
||||
@@ -171,19 +178,22 @@ class BaseModelService(ABC):
|
||||
ufs = self.settings.get("version_grouping", "same_base")
|
||||
group_by_base = ufs == "same_base"
|
||||
|
||||
model_groups: Dict[Any, List[Dict]] = {}
|
||||
ungrouped_standalone: List[Dict] = []
|
||||
model_groups: Dict[Any, List[Dict[str, Any]]] = {}
|
||||
ungrouped_standalone: List[Dict[str, Any]] = []
|
||||
for item in sorted_data:
|
||||
mid = self._extract_model_id(item)
|
||||
mid = self._extract_group_key(item)
|
||||
if mid is None:
|
||||
ungrouped_standalone.append(item)
|
||||
continue
|
||||
key = (mid, item.get("base_model") or "") if group_by_base else mid
|
||||
model_groups.setdefault(key, []).append(item)
|
||||
# Sort versions within each group by version id descending
|
||||
# Sort versions within each group by version id (descending);
|
||||
# fall back to modified timestamp for non-CivitAI sources.
|
||||
for items in model_groups.values():
|
||||
items.sort(
|
||||
key=lambda x: self._extract_version_id(x) or 0,
|
||||
key=lambda x: self._extract_version_id(x)
|
||||
or x.get("modified", 0)
|
||||
or 0,
|
||||
reverse=True,
|
||||
)
|
||||
# Sort groups by version count
|
||||
@@ -239,7 +249,7 @@ class BaseModelService(ABC):
|
||||
filter_duration = time.perf_counter() - t1
|
||||
post_filter_count = len(filtered_data)
|
||||
|
||||
annotated_for_filter: Optional[List[Dict]] = None
|
||||
annotated_for_filter: Optional[List[Dict[str, Any]]] = None
|
||||
t2 = time.perf_counter()
|
||||
if update_available_only:
|
||||
annotated_for_filter = await self._annotate_update_flags(filtered_data)
|
||||
@@ -286,11 +296,11 @@ class BaseModelService(ABC):
|
||||
page: int,
|
||||
page_size: int,
|
||||
sort_by: str = "name",
|
||||
search: str = None,
|
||||
search: str | None = None,
|
||||
fuzzy_search: bool = False,
|
||||
search_options: dict = None,
|
||||
search_options: dict[str, Any] | None = None,
|
||||
**kwargs,
|
||||
) -> Dict:
|
||||
) -> Dict[str, Any]:
|
||||
"""Get paginated excluded model data."""
|
||||
excluded_paths = list(self.scanner.get_excluded_models())
|
||||
excluded_entries: List[Dict[str, Any]] = []
|
||||
@@ -316,7 +326,7 @@ class BaseModelService(ABC):
|
||||
]
|
||||
persist_current_cache = getattr(self.scanner, "_persist_current_cache", None)
|
||||
if callable(persist_current_cache):
|
||||
await persist_current_cache()
|
||||
await cast(Awaitable[Any], persist_current_cache())
|
||||
|
||||
excluded_entries = self._sort_entries(excluded_entries, sort_by)
|
||||
|
||||
@@ -381,6 +391,12 @@ class BaseModelService(ABC):
|
||||
(item.get("model_name") or item.get("file_name") or "").lower(),
|
||||
item.get("file_path", "").lower(),
|
||||
)
|
||||
elif key_name == "random":
|
||||
# Seeded random shuffle: same seed -> same order (stable pagination)
|
||||
rng = random.Random(sort_params.seed or "random")
|
||||
result = list(data)
|
||||
rng.shuffle(result)
|
||||
return result
|
||||
elif key_name == "size":
|
||||
key_fn = lambda item: (
|
||||
int(item.get("size", 0) or 0),
|
||||
@@ -428,39 +444,50 @@ class BaseModelService(ABC):
|
||||
return entry
|
||||
|
||||
async def _apply_hash_filters(
|
||||
self, data: List[Dict], hash_filters: Dict
|
||||
) -> List[Dict]:
|
||||
"""Apply hash-based filtering"""
|
||||
self, data: List[Dict[str, Any]], hash_filters: Dict[str, Any]
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Apply hash-based filtering (SHA256 and AutoV3)."""
|
||||
|
||||
def matches_hash_set(item: Dict[str, Any], hash_set: set[str]) -> bool:
|
||||
"""Check whether an item matches any hash in the set.
|
||||
|
||||
Compares the item's ``sha256`` field and its non-empty ``autov3``
|
||||
field, both case-insensitively.
|
||||
"""
|
||||
if item.get("sha256", "").lower() in hash_set:
|
||||
return True
|
||||
autov3 = item.get("autov3", "")
|
||||
return bool(autov3) and autov3.lower() in hash_set
|
||||
|
||||
single_hash = hash_filters.get("single_hash")
|
||||
multiple_hashes = hash_filters.get("multiple_hashes")
|
||||
|
||||
if single_hash:
|
||||
# Filter by single hash
|
||||
single_hash = single_hash.lower()
|
||||
# Filter by single hash (SHA256 or AutoV3)
|
||||
return [
|
||||
item for item in data if item.get("sha256", "").lower() == single_hash
|
||||
item for item in data if matches_hash_set(item, {single_hash.lower()})
|
||||
]
|
||||
elif multiple_hashes:
|
||||
# Filter by multiple hashes
|
||||
hash_set = set(hash.lower() for hash in multiple_hashes)
|
||||
return [item for item in data if item.get("sha256", "").lower() in hash_set]
|
||||
# Filter by multiple hashes (SHA256 or AutoV3)
|
||||
hash_set = {hash.lower() for hash in multiple_hashes}
|
||||
return [item for item in data if matches_hash_set(item, hash_set)]
|
||||
|
||||
return data
|
||||
|
||||
async def _apply_common_filters(
|
||||
self,
|
||||
data: List[Dict],
|
||||
folder: str = None,
|
||||
folder_include: list = None,
|
||||
folder_exclude: list = None,
|
||||
base_models: list = None,
|
||||
model_types: list = None,
|
||||
data: List[Dict[str, Any]],
|
||||
folder: str | None = None,
|
||||
folder_include: list[str] | None = None,
|
||||
folder_exclude: list[str] | None = None,
|
||||
base_models: list[str] | None = None,
|
||||
model_types: list[str] | None = None,
|
||||
tags: Optional[Dict[str, str]] = None,
|
||||
auto_tags: Optional[Dict[str, str]] = None,
|
||||
favorites_only: bool = False,
|
||||
search_options: dict = None,
|
||||
search_options: dict[str, Any] | None = None,
|
||||
tag_logic: str = "any",
|
||||
) -> List[Dict]:
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Apply common filters that work across all model types"""
|
||||
normalized_options = self.search_strategy.normalize_options(search_options)
|
||||
criteria = FilterCriteria(
|
||||
@@ -479,24 +506,24 @@ class BaseModelService(ABC):
|
||||
|
||||
async def _apply_search_filters(
|
||||
self,
|
||||
data: List[Dict],
|
||||
data: List[Dict[str, Any]],
|
||||
search: str,
|
||||
fuzzy_search: bool,
|
||||
search_options: dict,
|
||||
) -> List[Dict]:
|
||||
search_options: dict[str, Any] | None,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Apply search filtering"""
|
||||
normalized_options = self.search_strategy.normalize_options(search_options)
|
||||
return self.search_strategy.apply(
|
||||
data, search, normalized_options, fuzzy_search
|
||||
)
|
||||
|
||||
async def _apply_specific_filters(self, data: List[Dict], **kwargs) -> List[Dict]:
|
||||
async def _apply_specific_filters(self, data: List[Dict[str, Any]], **kwargs) -> List[Dict[str, Any]]:
|
||||
"""Apply model-specific filters - to be overridden by subclasses if needed"""
|
||||
return data
|
||||
|
||||
async def _apply_credit_required_filter(
|
||||
self, data: List[Dict], credit_required: bool
|
||||
) -> List[Dict]:
|
||||
self, data: List[Dict[str, Any]], credit_required: bool
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Apply credit required filtering based on license_flags.
|
||||
|
||||
Args:
|
||||
@@ -526,8 +553,8 @@ class BaseModelService(ABC):
|
||||
return filtered_data
|
||||
|
||||
async def _apply_allow_selling_filter(
|
||||
self, data: List[Dict], allow_selling: bool
|
||||
) -> List[Dict]:
|
||||
self, data: List[Dict[str, Any]], allow_selling: bool
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Apply allow selling generated content filtering based on license_flags.
|
||||
|
||||
Args:
|
||||
@@ -559,8 +586,8 @@ class BaseModelService(ABC):
|
||||
|
||||
async def _annotate_update_flags(
|
||||
self,
|
||||
items: List[Dict],
|
||||
) -> List[Dict]:
|
||||
items: List[Dict[str, Any]],
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Attach an update_available flag to each response item.
|
||||
|
||||
Items without a civitai model id default to False.
|
||||
@@ -575,7 +602,7 @@ class BaseModelService(ABC):
|
||||
item["update_available"] = False
|
||||
return annotated
|
||||
|
||||
id_to_items: Dict[int, List[Dict]] = {}
|
||||
id_to_items: Dict[int, List[Dict[str, Any]]] = {}
|
||||
ordered_ids: List[int] = []
|
||||
for item in annotated:
|
||||
model_id = self._extract_model_id(item)
|
||||
@@ -606,15 +633,25 @@ class BaseModelService(ABC):
|
||||
except Exception:
|
||||
hide_early_access = False
|
||||
|
||||
# Check user setting for hiding permanent paid updates
|
||||
hide_paid = False
|
||||
try:
|
||||
hide_paid = bool(self.settings.get("hide_paid_updates", False))
|
||||
except Exception:
|
||||
hide_paid = False
|
||||
|
||||
records = None
|
||||
resolved: Optional[Dict[int, bool]] = None
|
||||
if same_base_mode:
|
||||
record_method = getattr(self.update_service, "get_records_bulk", None)
|
||||
if callable(record_method):
|
||||
try:
|
||||
records = await record_method(self.model_type, ordered_ids)
|
||||
records = await cast(Awaitable[Any], record_method(self.model_type, ordered_ids))
|
||||
resolved = {
|
||||
model_id: record.has_update(hide_early_access=hide_early_access)
|
||||
model_id: record.has_update(
|
||||
hide_early_access=hide_early_access,
|
||||
hide_paid=hide_paid,
|
||||
)
|
||||
for model_id, record in records.items()
|
||||
}
|
||||
except Exception as exc:
|
||||
@@ -632,11 +669,12 @@ class BaseModelService(ABC):
|
||||
bulk_method = getattr(self.update_service, "has_updates_bulk", None)
|
||||
if callable(bulk_method):
|
||||
try:
|
||||
resolved = await bulk_method(
|
||||
resolved = await cast(Awaitable[Any], bulk_method(
|
||||
self.model_type,
|
||||
ordered_ids,
|
||||
hide_early_access=hide_early_access,
|
||||
)
|
||||
hide_paid=hide_paid,
|
||||
))
|
||||
except Exception as exc:
|
||||
logger.error(
|
||||
"Failed to resolve update status in bulk for %s models (%s): %s",
|
||||
@@ -650,7 +688,10 @@ class BaseModelService(ABC):
|
||||
if resolved is None:
|
||||
tasks = [
|
||||
self.update_service.has_update(
|
||||
self.model_type, model_id, hide_early_access=hide_early_access
|
||||
self.model_type,
|
||||
model_id,
|
||||
hide_early_access=hide_early_access,
|
||||
hide_paid=hide_paid,
|
||||
)
|
||||
for model_id in ordered_ids
|
||||
]
|
||||
@@ -690,6 +731,7 @@ class BaseModelService(ABC):
|
||||
threshold_version,
|
||||
base_model,
|
||||
hide_early_access=hide_early_access,
|
||||
hide_paid=hide_paid,
|
||||
)
|
||||
else:
|
||||
flag = default_flag
|
||||
@@ -698,7 +740,34 @@ class BaseModelService(ABC):
|
||||
return annotated
|
||||
|
||||
@staticmethod
|
||||
def _extract_model_id(item: Dict) -> Optional[int]:
|
||||
def _extract_hf_group_key(item: Dict[str, Any]) -> Optional[str]:
|
||||
"""Extract `hf:{owner}/{repo}` from item's ``hf_url``, or None."""
|
||||
hf_url = item.get("hf_url") if isinstance(item, dict) else None
|
||||
if not hf_url or not isinstance(hf_url, str):
|
||||
return None
|
||||
m = re.match(
|
||||
r"https?://huggingface\.co/([^/]+/[^/]+)", hf_url.strip()
|
||||
)
|
||||
if not m:
|
||||
return None
|
||||
return f"hf:{m.group(1)}"
|
||||
|
||||
@staticmethod
|
||||
def _extract_group_key(item: Dict[str, Any]) -> Union[int, str, None]:
|
||||
"""Return the group identity key: CivitAI modelId (int) or HF repo (str).
|
||||
|
||||
Preference order:
|
||||
1. CivitAI ``modelId`` (int)
|
||||
2. HF repo identity ``hf:{owner}/{repo}`` (str)
|
||||
3. ``None`` (no known grouping source)
|
||||
"""
|
||||
mid = BaseModelService._extract_model_id(item)
|
||||
if mid is not None:
|
||||
return mid
|
||||
return BaseModelService._extract_hf_group_key(item)
|
||||
|
||||
@staticmethod
|
||||
def _extract_model_id(item: Dict[str, Any]) -> Optional[int]:
|
||||
civitai = item.get("civitai") if isinstance(item, dict) else None
|
||||
if not isinstance(civitai, dict):
|
||||
return None
|
||||
@@ -711,7 +780,7 @@ class BaseModelService(ABC):
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _extract_version_id(item: Dict) -> Optional[int]:
|
||||
def _extract_version_id(item: Dict[str, Any]) -> Optional[int]:
|
||||
civitai = item.get("civitai") if isinstance(item, dict) else None
|
||||
if not isinstance(civitai, dict):
|
||||
return None
|
||||
@@ -724,7 +793,7 @@ class BaseModelService(ABC):
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _extract_base_model(item: Dict) -> Optional[str]:
|
||||
def _extract_base_model(item: Dict[str, Any]) -> Optional[str]:
|
||||
value = item.get("base_model")
|
||||
if value is None:
|
||||
return None
|
||||
@@ -776,7 +845,7 @@ class BaseModelService(ABC):
|
||||
|
||||
return highest_by_base
|
||||
|
||||
def _paginate(self, data: List[Dict], page: int, page_size: int) -> Dict:
|
||||
def _paginate(self, data: List[Dict[str, Any]], page: int, page_size: int) -> Dict[str, Any]:
|
||||
"""Apply pagination to filtered data"""
|
||||
total_items = len(data)
|
||||
start_idx = (page - 1) * page_size
|
||||
@@ -791,7 +860,7 @@ class BaseModelService(ABC):
|
||||
}
|
||||
|
||||
@abstractmethod
|
||||
async def format_response(self, model_data: Dict) -> Optional[Dict]:
|
||||
async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||
"""Format model data for API response - must be implemented by subclasses.
|
||||
|
||||
Subclasses should return None for corrupted entries so the handler
|
||||
@@ -800,11 +869,17 @@ class BaseModelService(ABC):
|
||||
pass
|
||||
|
||||
# Common service methods that delegate to scanner
|
||||
async def get_top_tags(self, limit: int = 20) -> List[Dict]:
|
||||
async def get_top_tags(self, limit: int = 20) -> List[Dict[str, Any]]:
|
||||
"""Get top tags sorted by frequency"""
|
||||
return await self.scanner.get_top_tags(limit)
|
||||
|
||||
async def get_base_models(self, limit: int = 20) -> List[Dict]:
|
||||
async def search_tags(
|
||||
self, query: str, limit: int = 50
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Search tags by substring, sorted by frequency"""
|
||||
return await self.scanner.search_tags(query, limit)
|
||||
|
||||
async def get_base_models(self, limit: int = 20) -> List[Dict[str, Any]]:
|
||||
"""Get base models sorted by frequency"""
|
||||
return await self.scanner.get_base_models(limit)
|
||||
|
||||
@@ -871,7 +946,7 @@ class BaseModelService(ABC):
|
||||
"""Get model root directories"""
|
||||
return self.scanner.get_model_roots()
|
||||
|
||||
def filter_civitai_data(self, data: Dict, minimal: bool = False) -> Dict:
|
||||
def filter_civitai_data(self, data: Dict[str, Any], minimal: bool = False) -> Dict[str, Any]:
|
||||
"""Filter relevant fields from CivitAI data"""
|
||||
if not data:
|
||||
return {}
|
||||
@@ -897,14 +972,25 @@ class BaseModelService(ABC):
|
||||
)
|
||||
return {k: data[k] for k in fields if k in data}
|
||||
|
||||
async def get_folder_tree(self, model_root: str) -> Dict:
|
||||
async def _get_tree_folders(self, cache, include_empty: bool) -> List[str]:
|
||||
"""Return the folder list backing folder tree responses.
|
||||
|
||||
With ``include_empty`` the directories are enumerated live from the
|
||||
filesystem (including empty ones) via the scanner; otherwise the
|
||||
models-only ``cache.folders`` list is used unchanged.
|
||||
"""
|
||||
if include_empty:
|
||||
return await self.scanner.get_all_folders()
|
||||
return cache.folders
|
||||
|
||||
async def get_folder_tree(self, model_root: str, include_empty: bool = False) -> Dict[str, Any]:
|
||||
"""Get hierarchical folder tree for a specific model root"""
|
||||
cache = await self.scanner.get_cached_data()
|
||||
|
||||
# Build tree structure from folders
|
||||
tree = {}
|
||||
|
||||
for folder in cache.folders:
|
||||
for folder in await self._get_tree_folders(cache, include_empty):
|
||||
# Check if this folder belongs to the specified model root
|
||||
folder_belongs_to_root = False
|
||||
for root in self.scanner.get_model_roots():
|
||||
@@ -926,7 +1012,7 @@ class BaseModelService(ABC):
|
||||
|
||||
return tree
|
||||
|
||||
async def get_unified_folder_tree(self) -> Dict:
|
||||
async def get_unified_folder_tree(self, include_empty: bool = False) -> Dict[str, Any]:
|
||||
"""Get unified folder tree across all model roots"""
|
||||
cache = await self.scanner.get_cached_data()
|
||||
|
||||
@@ -936,7 +1022,7 @@ class BaseModelService(ABC):
|
||||
# Get all model roots for path normalization
|
||||
model_roots = self.scanner.get_model_roots()
|
||||
|
||||
for folder in cache.folders:
|
||||
for folder in await self._get_tree_folders(cache, include_empty):
|
||||
if not folder: # Skip empty folders
|
||||
continue
|
||||
|
||||
@@ -955,13 +1041,21 @@ class BaseModelService(ABC):
|
||||
|
||||
return unified_tree
|
||||
|
||||
async def get_model_notes(self, model_name: str) -> Optional[str]:
|
||||
"""Get notes for a specific model file"""
|
||||
async def get_model_notes(self, model_name: str) -> Optional[dict[str, Any]]:
|
||||
"""Get notes and file_path for a specific model file.
|
||||
|
||||
Supports both simple names (``OWSMianne_ANIMA_V1``) and full-path
|
||||
syntax (``Anima/character/OWSMianne_ANIMA_V1``).
|
||||
"""
|
||||
cache = await self.scanner.get_cached_data()
|
||||
|
||||
for model in cache.raw_data:
|
||||
if model["file_name"] == model_name:
|
||||
return model.get("notes", "")
|
||||
file_name = model.get("file_name", "")
|
||||
if file_name == model_name or model_name.endswith("/" + file_name) or model_name.endswith("\\" + file_name):
|
||||
return {
|
||||
"notes": model.get("notes", ""),
|
||||
"file_path": model.get("file_path", ""),
|
||||
}
|
||||
|
||||
return None
|
||||
|
||||
@@ -1079,11 +1173,16 @@ class BaseModelService(ABC):
|
||||
|
||||
return {"civitai_url": None, "model_id": None, "version_id": None}
|
||||
|
||||
async def get_model_metadata(self, file_path: str) -> Optional[Dict]:
|
||||
async def get_model_metadata(self, file_path: str) -> Optional[Dict[str, Any]]:
|
||||
"""Load full metadata for a single model.
|
||||
|
||||
Listing/search endpoints return lightweight cache entries; this method performs
|
||||
a lazy read of the on-disk metadata snapshot when callers need full detail.
|
||||
|
||||
As a beneficial side effect, the in-memory and persistent caches are
|
||||
opportunistically synchronised with the on-disk metadata — this keeps the
|
||||
caches fresh even when a ``.metadata.json`` file was edited outside of the
|
||||
normal save path (e.g. manually or by an external script).
|
||||
"""
|
||||
metadata, should_skip = await MetadataManager.load_metadata(
|
||||
file_path, self.metadata_class
|
||||
@@ -1101,6 +1200,19 @@ class BaseModelService(ABC):
|
||||
MetadataManager.save_metadata(file_path, metadata)
|
||||
)
|
||||
|
||||
# Opportunistically sync the in-memory + persistent caches.
|
||||
# The .metadata.json disk read is already paid for; the sync only
|
||||
# performs work when the cache is actually stale, and uses targeted,
|
||||
# in-place operations to minimise overhead even with large model sets.
|
||||
#
|
||||
# Fire-and-forget by design: the task is intentionally untracked.
|
||||
# sync_cache_from_metadata handles its own errors internally.
|
||||
asyncio.create_task(
|
||||
self.scanner.sync_cache_from_metadata(
|
||||
file_path, metadata.to_dict()
|
||||
)
|
||||
)
|
||||
|
||||
return self.filter_civitai_data(metadata.to_dict().get("civitai", {}))
|
||||
|
||||
async def get_model_description(self, file_path: str) -> Optional[str]:
|
||||
@@ -1157,7 +1269,7 @@ class BaseModelService(ABC):
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _relative_path_sort_key(relative_path: str, include_terms: List[str]) -> tuple:
|
||||
def _relative_path_sort_key(relative_path: str, include_terms: List[str]) -> tuple[int, int, int, str]:
|
||||
"""Sort paths by how well they satisfy the include tokens.
|
||||
|
||||
Sorts based on path without extension for consistent ordering.
|
||||
@@ -1183,20 +1295,109 @@ class BaseModelService(ABC):
|
||||
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(
|
||||
self, search_term: str, limit: int = 15, offset: int = 0
|
||||
self,
|
||||
search_term: str,
|
||||
limit: int = 15,
|
||||
offset: int = 0,
|
||||
*,
|
||||
folder: Optional[str] = None,
|
||||
folder_include: Optional[list[str]] = None,
|
||||
folder_exclude: Optional[list[str]] = None,
|
||||
base_models: Optional[list[str]] = None,
|
||||
model_types: Optional[list[str]] = None,
|
||||
tags: Optional[dict[str, str]] = None,
|
||||
auto_tags: Optional[dict[str, str]] = None,
|
||||
tag_logic: str = "any",
|
||||
credit_required: Optional[bool] = None,
|
||||
allow_selling_generated_content: Optional[bool] = None,
|
||||
recursive: bool = True,
|
||||
apply_filters: bool = False,
|
||||
) -> List[str]:
|
||||
"""Search model relative file paths for autocomplete functionality"""
|
||||
"""Search model relative file paths for autocomplete functionality.
|
||||
|
||||
Optional filter kwargs mirror the filters used by the list endpoint
|
||||
(/api/lm/{prefix}/list). When no filter kwargs are provided the
|
||||
behavior is identical to plain token-based path matching.
|
||||
"""
|
||||
cache = await self.scanner.get_cached_data()
|
||||
include_terms, exclude_terms = self._parse_search_tokens(search_term)
|
||||
|
||||
data = cache.raw_data
|
||||
has_filters = any(
|
||||
[
|
||||
apply_filters,
|
||||
folder is not None,
|
||||
folder_include,
|
||||
folder_exclude,
|
||||
base_models,
|
||||
model_types,
|
||||
tags,
|
||||
auto_tags,
|
||||
credit_required is not None,
|
||||
allow_selling_generated_content is not None,
|
||||
]
|
||||
)
|
||||
if has_filters:
|
||||
# Auto-tags are not stored in the scanner cache — they are computed
|
||||
# on the fly. Pre-compute them only when an auto-tag filter is
|
||||
# active to avoid mutating cache entries unnecessarily.
|
||||
if auto_tags:
|
||||
from .auto_tag_service import extract_auto_tags
|
||||
|
||||
for item in data:
|
||||
if not item.get("auto_tags"):
|
||||
item["auto_tags"] = extract_auto_tags(item)
|
||||
|
||||
criteria = FilterCriteria(
|
||||
folder=folder,
|
||||
folder_include=folder_include,
|
||||
folder_exclude=folder_exclude,
|
||||
base_models=base_models,
|
||||
model_types=model_types,
|
||||
tags=tags,
|
||||
auto_tags=auto_tags,
|
||||
search_options={"recursive": recursive},
|
||||
tag_logic=tag_logic,
|
||||
)
|
||||
data = self.filter_set.apply(data, criteria)
|
||||
if credit_required is not None:
|
||||
data = await self._apply_credit_required_filter(
|
||||
data, credit_required
|
||||
)
|
||||
if allow_selling_generated_content is not None:
|
||||
data = await self._apply_allow_selling_filter(
|
||||
data, allow_selling_generated_content
|
||||
)
|
||||
|
||||
matching_paths = []
|
||||
|
||||
# Get model roots for path calculation
|
||||
model_roots = self.scanner.get_model_roots()
|
||||
|
||||
# Collect all matching paths first (needed for proper sorting and offset)
|
||||
for model in cache.raw_data:
|
||||
for model in data:
|
||||
file_path = model.get("file_path", "")
|
||||
if not file_path:
|
||||
continue
|
||||
@@ -1224,9 +1425,13 @@ class BaseModelService(ABC):
|
||||
):
|
||||
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(
|
||||
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
|
||||
|
||||
@@ -20,6 +20,11 @@ from .recipes import (
|
||||
RecipeDownloadError,
|
||||
RecipeNotFoundError,
|
||||
)
|
||||
from .recipes.import_info import (
|
||||
CHANNEL_BATCH_IMPORT_LOCAL,
|
||||
CHANNEL_BATCH_IMPORT_URL,
|
||||
build_import_info,
|
||||
)
|
||||
|
||||
|
||||
class ImportItemType(Enum):
|
||||
@@ -71,6 +76,9 @@ class BatchImportProgress:
|
||||
tags: List[str] = field(default_factory=list)
|
||||
skip_no_metadata: 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]:
|
||||
return {
|
||||
@@ -82,6 +90,7 @@ class BatchImportProgress:
|
||||
"skipped": self.skipped,
|
||||
"current_item": self.current_item,
|
||||
"status": self.status,
|
||||
"rate_limited": self.rate_limited,
|
||||
"started_at": self.started_at,
|
||||
"finished_at": self.finished_at,
|
||||
"progress_percent": round((self.completed / self.total) * 100, 1)
|
||||
@@ -118,6 +127,10 @@ class AdaptiveConcurrencyController:
|
||||
self._task_durations: List[float] = []
|
||||
self._recent_errors = 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:
|
||||
self._task_durations.append(duration)
|
||||
@@ -146,7 +159,37 @@ class AdaptiveConcurrencyController:
|
||||
self._recent_successes = 0
|
||||
|
||||
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:
|
||||
@@ -184,6 +227,7 @@ class BatchImportService:
|
||||
def cancel_import(self, operation_id: str) -> bool:
|
||||
if operation_id in self._active_operations:
|
||||
self._cancellation_flags[operation_id] = True
|
||||
self._logger.info("Cancel requested for batch import operation %s", operation_id)
|
||||
return True
|
||||
return False
|
||||
|
||||
@@ -273,6 +317,14 @@ class BatchImportService:
|
||||
self._active_operations[operation_id] = progress
|
||||
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(
|
||||
self._run_batch_import(
|
||||
operation_id=operation_id,
|
||||
@@ -295,6 +347,12 @@ class BatchImportService:
|
||||
skip_duplicates: bool = False,
|
||||
) -> str:
|
||||
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]
|
||||
|
||||
@@ -334,6 +392,13 @@ class BatchImportService:
|
||||
ext = os.path.splitext(filename)[1].lower()
|
||||
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(
|
||||
self,
|
||||
*,
|
||||
@@ -379,6 +444,9 @@ class BatchImportService:
|
||||
self._concurrency_controller.record_result(
|
||||
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"):
|
||||
item.status = ImportStatus.SUCCESS
|
||||
@@ -389,6 +457,17 @@ class BatchImportService:
|
||||
item.status = ImportStatus.SKIPPED
|
||||
item.error_message = result.get("error")
|
||||
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:
|
||||
item.status = ImportStatus.FAILED
|
||||
item.error_message = result.get("error")
|
||||
@@ -396,13 +475,36 @@ class BatchImportService:
|
||||
|
||||
except Exception as e:
|
||||
self._logger.error(f"Error importing {item.source}: {e}")
|
||||
item.duration = time.time() - start_time
|
||||
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)
|
||||
item.duration = time.time() - start_time
|
||||
progress.failed += 1
|
||||
self._concurrency_controller.record_result(item.duration, False)
|
||||
await self._concurrency_controller.apply_concurrency()
|
||||
|
||||
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)
|
||||
|
||||
tasks = [process_item(item) for item in progress.items]
|
||||
@@ -415,6 +517,15 @@ class BatchImportService:
|
||||
|
||||
progress.finished_at = time.time()
|
||||
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 asyncio.sleep(5)
|
||||
@@ -518,6 +629,17 @@ class BatchImportService:
|
||||
"loras": loras,
|
||||
"gen_params": payload.get("gen_params", {}),
|
||||
"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"):
|
||||
@@ -595,3 +717,6 @@ class BatchImportService:
|
||||
def _cleanup_operation(self, operation_id: str) -> None:
|
||||
if operation_id in self._cancellation_flags:
|
||||
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)
|
||||
|
||||
@@ -59,6 +59,7 @@ class CacheEntryValidator:
|
||||
'notes': ('', False),
|
||||
'usage_tips': ('', False),
|
||||
'hash_status': ('completed', False),
|
||||
'autov3': (None, False),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
@@ -119,6 +120,11 @@ class CacheEntryValidator:
|
||||
if is_required:
|
||||
errors.append(f"Required field '{field_name}' is missing or None")
|
||||
if auto_repair:
|
||||
# A missing optional field whose default is None is already
|
||||
# semantically equal to its default (e.g. autov3: absent
|
||||
# means "not checked") — writing None back is a no-op, not
|
||||
# a repair.
|
||||
if default_value is not None:
|
||||
working_entry[field_name] = cls._get_default_copy(default_value)
|
||||
repaired = True
|
||||
continue
|
||||
@@ -175,6 +181,15 @@ class CacheEntryValidator:
|
||||
# that invalidates the entry, but we also don't mark it repaired.
|
||||
pass
|
||||
|
||||
# Normalize autov3 to lowercase if needed (optional field, never stripped).
|
||||
autov3 = working_entry.get('autov3')
|
||||
if isinstance(autov3, str) and autov3:
|
||||
normalized_autov3 = autov3.lower()
|
||||
if normalized_autov3 != autov3:
|
||||
if auto_repair:
|
||||
working_entry['autov3'] = normalized_autov3
|
||||
repaired = True
|
||||
|
||||
# Determine if entry is valid
|
||||
# Entry is valid if no critical required field errors remain after repair
|
||||
# Critical fields are file_path and sha256
|
||||
@@ -242,6 +257,19 @@ class CacheEntryValidator:
|
||||
"""
|
||||
expected_type = type(default_value)
|
||||
|
||||
# Special case: autov3 is optional with a three-state contract.
|
||||
# None = not checked, "" = checked but unavailable, otherwise a
|
||||
# 12-character hex string (case-insensitive here; normalized to
|
||||
# lowercase separately).
|
||||
if field_name == 'autov3':
|
||||
if value is None or value == "":
|
||||
return None
|
||||
if not isinstance(value, str):
|
||||
return f"Field 'autov3' should be string or None, got {type(value).__name__}"
|
||||
if len(value) != 12 or any(c not in '0123456789abcdefABCDEF' for c in value):
|
||||
return "Field 'autov3' should be a 12-character hex string"
|
||||
return None
|
||||
|
||||
# Special handling for numeric types
|
||||
if expected_type == int:
|
||||
if not isinstance(value, (int, float)):
|
||||
|
||||
@@ -1,3 +1,7 @@
|
||||
# pyright: reportImportCycles=false
|
||||
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||
# import cycles. Breaking them would require an architectural refactor.
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
@@ -6,10 +10,10 @@ from datetime import datetime
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from ..utils.models import CheckpointMetadata
|
||||
from ..utils.file_utils import find_preview_file, normalize_path
|
||||
from ..utils.file_utils import find_preview_file, normalize_path, calculate_autov3
|
||||
from ..utils.metadata_manager import MetadataManager
|
||||
from ..config import config
|
||||
from .model_scanner import ModelScanner
|
||||
from .model_scanner import ModelScanner, _is_excluded_dir
|
||||
from .model_hash_index import ModelHashIndex
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -62,6 +66,11 @@ class CheckpointScanner(ModelScanner):
|
||||
# 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 checkpoints; record the checked state at creation time ("" =
|
||||
# checked but unavailable).
|
||||
autov3 = calculate_autov3(real_path)
|
||||
|
||||
# Create metadata WITHOUT calculating hash
|
||||
metadata = CheckpointMetadata(
|
||||
file_name=base_name,
|
||||
@@ -77,6 +86,7 @@ class CheckpointScanner(ModelScanner):
|
||||
sub_type="checkpoint",
|
||||
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
|
||||
@@ -114,6 +124,17 @@ class CheckpointScanner(ModelScanner):
|
||||
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:
|
||||
@@ -125,6 +146,11 @@ class CheckpointScanner(ModelScanner):
|
||||
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)
|
||||
@@ -175,6 +201,13 @@ class CheckpointScanner(ModelScanner):
|
||||
|
||||
# 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
|
||||
@@ -191,7 +224,26 @@ class CheckpointScanner(ModelScanner):
|
||||
await MetadataManager.save_metadata(file_path, metadata)
|
||||
|
||||
# Update hash index
|
||||
self._hash_index.add_entry(sha256.lower(), file_path)
|
||||
self._hash_index.add_entry(
|
||||
sha256.lower(),
|
||||
file_path,
|
||||
getattr(metadata, "autov3", None) or None,
|
||||
)
|
||||
|
||||
# Update the in-memory cache entry so that subsequent
|
||||
# _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 checkpoint: {file_path}")
|
||||
return sha256
|
||||
@@ -276,7 +328,8 @@ class CheckpointScanner(ModelScanner):
|
||||
if not os.path.exists(root_path):
|
||||
continue
|
||||
|
||||
for dirpath, _dirnames, filenames in os.walk(root_path):
|
||||
for dirpath, dirnames, filenames in os.walk(root_path):
|
||||
dirnames[:] = [d for d in dirnames if not _is_excluded_dir(d)]
|
||||
for filename in filenames:
|
||||
if not filename.endswith(".metadata.json"):
|
||||
continue
|
||||
@@ -380,7 +433,7 @@ class CheckpointScanner(ModelScanner):
|
||||
roots.extend(config.extra_checkpoints_roots or [])
|
||||
roots.extend(config.extra_unet_roots or [])
|
||||
# Remove duplicates while preserving order
|
||||
seen: set = set()
|
||||
seen: set[str] = set()
|
||||
unique_roots: List[str] = []
|
||||
for root in roots:
|
||||
if root not in seen:
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import os
|
||||
import logging
|
||||
from typing import Dict, Optional
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from .base_model_service import BaseModelService
|
||||
from .auto_tag_service import extract_auto_tags
|
||||
@@ -21,58 +21,59 @@ class CheckpointService(BaseModelService):
|
||||
"""
|
||||
super().__init__("checkpoint", scanner, CheckpointMetadata, update_service=update_service)
|
||||
|
||||
async def format_response(self, checkpoint_data: Dict) -> Optional[Dict]:
|
||||
async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||
"""Format Checkpoint data for API response.
|
||||
|
||||
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 = checkpoint_data.get("file_path")
|
||||
file_path = model_data.get("file_path")
|
||||
if not file_path or not isinstance(file_path, str):
|
||||
logger.warning(
|
||||
"Skipping corrupted checkpoint entry (missing file_path): %s",
|
||||
checkpoint_data.get("file_name", "<unknown>"),
|
||||
model_data.get("file_name", "<unknown>"),
|
||||
)
|
||||
return None
|
||||
|
||||
# Get sub_type from cache entry (new canonical field)
|
||||
sub_type = checkpoint_data.get("sub_type", "checkpoint")
|
||||
sub_type = model_data.get("sub_type", "checkpoint")
|
||||
|
||||
file_name = checkpoint_data.get("file_name") or ""
|
||||
model_name = checkpoint_data.get("model_name") or file_name
|
||||
folder = checkpoint_data.get("folder") or ""
|
||||
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(checkpoint_data.get("preview_url", "")),
|
||||
"preview_nsfw_level": checkpoint_data.get("preview_nsfw_level", 0),
|
||||
"base_model": checkpoint_data.get("base_model", ""),
|
||||
"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": checkpoint_data.get("sha256", ""),
|
||||
"sha256": model_data.get("sha256", ""),
|
||||
"autov3": model_data.get("autov3"),
|
||||
"file_path": file_path.replace(os.sep, "/"),
|
||||
"file_size": checkpoint_data.get("size", 0),
|
||||
"modified": checkpoint_data.get("modified", ""),
|
||||
"tags": checkpoint_data.get("tags", []),
|
||||
"from_civitai": checkpoint_data.get("from_civitai", True),
|
||||
"usage_count": checkpoint_data.get("usage_count", 0),
|
||||
"notes": checkpoint_data.get("notes", ""),
|
||||
"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),
|
||||
"usage_count": model_data.get("usage_count", 0),
|
||||
"notes": model_data.get("notes", ""),
|
||||
"sub_type": sub_type,
|
||||
"favorite": checkpoint_data.get("favorite", False),
|
||||
"exclude": bool(checkpoint_data.get("exclude", False)),
|
||||
"update_available": bool(checkpoint_data.get("update_available", False)),
|
||||
"skip_metadata_refresh": bool(checkpoint_data.get("skip_metadata_refresh", False)),
|
||||
"civitai": self.filter_civitai_data(checkpoint_data.get("civitai", {}), minimal=True),
|
||||
"auto_tags": checkpoint_data.get("auto_tags") or extract_auto_tags(checkpoint_data),
|
||||
"version_count": checkpoint_data.get("version_count"),
|
||||
"hf_url": checkpoint_data.get("hf_url", ""),
|
||||
"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"),
|
||||
"hf_url": model_data.get("hf_url", ""),
|
||||
}
|
||||
|
||||
def find_duplicate_hashes(self) -> Dict:
|
||||
def find_duplicate_hashes(self) -> Dict[str, Any]:
|
||||
"""Find Checkpoints with duplicate SHA256 hashes"""
|
||||
return self.scanner._hash_index.get_duplicate_hashes()
|
||||
|
||||
def find_duplicate_filenames(self) -> Dict:
|
||||
def find_duplicate_filenames(self) -> Dict[str, Any]:
|
||||
"""Find Checkpoints with conflicting filenames"""
|
||||
return self.scanner._hash_index.get_duplicate_filenames()
|
||||
|
||||
@@ -1,8 +1,13 @@
|
||||
# pyright: reportImportCycles=false
|
||||
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||
# import cycles. Breaking them would require an architectural refactor.
|
||||
import json
|
||||
import logging
|
||||
import asyncio
|
||||
from copy import deepcopy
|
||||
from typing import Optional, Dict, Tuple, List
|
||||
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 .downloader import get_downloader
|
||||
from .errors import RateLimitError
|
||||
@@ -37,12 +42,16 @@ class CivArchiveClient:
|
||||
async def _request_json(
|
||||
self,
|
||||
path: str,
|
||||
params: Optional[Dict[str, str]] = None
|
||||
) -> Tuple[Optional[Dict], Optional[str]]:
|
||||
params: Optional[Dict[str, Any]] = None
|
||||
) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||
"""Call CivArchive API and return JSON payload"""
|
||||
success, payload = await self._make_request(path, params=params)
|
||||
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
|
||||
if not isinstance(payload, dict):
|
||||
return None, "Invalid response structure"
|
||||
@@ -52,12 +61,12 @@ class CivArchiveClient:
|
||||
self,
|
||||
path: str,
|
||||
*,
|
||||
params: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[bool, Dict | str]:
|
||||
params: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[bool, Dict[str, Any] | str]:
|
||||
"""Wrapper around downloader.make_request that surfaces rate limits."""
|
||||
|
||||
downloader = await get_downloader()
|
||||
kwargs: Dict[str, Dict[str, str]] = {}
|
||||
kwargs: Dict[str, Dict[str, Any]] = {}
|
||||
if params:
|
||||
safe_params = {str(key): str(value) for key, value in params.items() if value is not None}
|
||||
if safe_params:
|
||||
@@ -73,10 +82,11 @@ class CivArchiveClient:
|
||||
if payload.provider is None:
|
||||
payload.provider = "civarchive_api"
|
||||
raise payload
|
||||
return success, payload
|
||||
# RateLimitError is always raised above, so the returned payload is a dict or str.
|
||||
return success, cast(Dict[str, Any] | str, payload)
|
||||
|
||||
@staticmethod
|
||||
def _normalize_payload(payload: Dict) -> Dict:
|
||||
def _normalize_payload(payload: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Unwrap CivArchive responses that wrap content under a data key"""
|
||||
if not isinstance(payload, dict):
|
||||
return {}
|
||||
@@ -86,12 +96,12 @@ class CivArchiveClient:
|
||||
return payload
|
||||
|
||||
@staticmethod
|
||||
def _split_context(payload: Dict) -> Tuple[Dict, Dict, List[Dict]]:
|
||||
def _split_context(payload: Dict[str, Any]) -> Tuple[Dict[str, Any], Dict[str, Any], List[Dict[str, Any]]]:
|
||||
"""Separate version payload from surrounding model context"""
|
||||
data = CivArchiveClient._normalize_payload(payload)
|
||||
context: Dict = {}
|
||||
fallback_files: List[Dict] = []
|
||||
version: Dict = {}
|
||||
context: Dict[str, Any] = {}
|
||||
fallback_files: List[Dict[str, Any]] = []
|
||||
version: Dict[str, Any] = {}
|
||||
|
||||
for key, value in data.items():
|
||||
if key in {"version", "model"}:
|
||||
@@ -115,7 +125,7 @@ class CivArchiveClient:
|
||||
return context, version, fallback_files
|
||||
|
||||
@staticmethod
|
||||
def _ensure_list(value) -> List:
|
||||
def _ensure_list(value: Any) -> List[Any]:
|
||||
if isinstance(value, list):
|
||||
return value
|
||||
if value is None:
|
||||
@@ -123,7 +133,7 @@ class CivArchiveClient:
|
||||
return [value]
|
||||
|
||||
@staticmethod
|
||||
def _build_model_info(context: Dict) -> Dict:
|
||||
def _build_model_info(context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
tags = context.get("tags")
|
||||
if not isinstance(tags, list):
|
||||
tags = list(tags) if isinstance(tags, (set, tuple)) else ([] if tags is None else [tags])
|
||||
@@ -136,7 +146,7 @@ class CivArchiveClient:
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _build_creator_info(context: Dict) -> Dict:
|
||||
def _build_creator_info(context: Dict[str, Any]) -> Dict[str, Any]:
|
||||
username = context.get("creator_username") or context.get("username") or ""
|
||||
image = context.get("creator_image") or context.get("creator_avatar") or ""
|
||||
creator: Dict[str, Optional[str]] = {
|
||||
@@ -150,7 +160,7 @@ class CivArchiveClient:
|
||||
return creator
|
||||
|
||||
@staticmethod
|
||||
def _transform_file_entry(file_data: Dict) -> Dict:
|
||||
def _transform_file_entry(file_data: Dict[str, Any]) -> Dict[str, Any]:
|
||||
mirrors = file_data.get("mirrors") or []
|
||||
if not isinstance(mirrors, list):
|
||||
mirrors = [mirrors]
|
||||
@@ -165,7 +175,7 @@ class CivArchiveClient:
|
||||
if not name and available_mirror:
|
||||
name = available_mirror.get("filename")
|
||||
|
||||
transformed: Dict = {
|
||||
transformed: Dict[str, Any] = {
|
||||
"id": file_data.get("id"),
|
||||
"sizeKB": file_data.get("sizeKB"),
|
||||
"name": name,
|
||||
@@ -216,23 +226,23 @@ class CivArchiveClient:
|
||||
|
||||
def _transform_files(
|
||||
self,
|
||||
files: Optional[List[Dict]],
|
||||
fallback_files: Optional[List[Dict]] = None
|
||||
) -> List[Dict]:
|
||||
candidates: List[Dict] = []
|
||||
files: Optional[List[Dict[str, Any]]],
|
||||
fallback_files: Optional[List[Dict[str, Any]]] = None
|
||||
) -> List[Dict[str, Any]]:
|
||||
candidates: List[Dict[str, Any]] = []
|
||||
if isinstance(files, list) and files:
|
||||
candidates = files
|
||||
elif isinstance(fallback_files, list):
|
||||
candidates = fallback_files
|
||||
|
||||
transformed_files: List[Dict] = []
|
||||
transformed_files: List[Dict[str, Any]] = []
|
||||
for file_data in candidates:
|
||||
if isinstance(file_data, dict):
|
||||
transformed_files.append(self._transform_file_entry(file_data))
|
||||
|
||||
# Sort: .safetensors first, .ckpt second, others last
|
||||
# so the backend fallback (no file_params) prefers safetensors
|
||||
def _sort_key(f: Dict) -> int:
|
||||
def _sort_key(f: Dict[str, Any]) -> int:
|
||||
fname = f.get("name") or ""
|
||||
if isinstance(fname, str):
|
||||
lower = fname.lower()
|
||||
@@ -247,10 +257,10 @@ class CivArchiveClient:
|
||||
|
||||
def _transform_version(
|
||||
self,
|
||||
context: Dict,
|
||||
version: Dict,
|
||||
fallback_files: Optional[List[Dict]] = None
|
||||
) -> Optional[Dict]:
|
||||
context: Dict[str, Any],
|
||||
version: Dict[str, Any],
|
||||
fallback_files: Optional[List[Dict[str, Any]]] = None
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
if not version:
|
||||
return None
|
||||
|
||||
@@ -291,8 +301,10 @@ class CivArchiveClient:
|
||||
|
||||
return version_copy
|
||||
|
||||
async def _resolve_version_from_files(self, payload: Dict) -> Optional[Dict]:
|
||||
async def _resolve_version_from_files(self, payload: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||
"""Fallback to fetch version data when only file metadata is available"""
|
||||
if not isinstance(payload, dict):
|
||||
return None
|
||||
data = self._normalize_payload(payload)
|
||||
files = data.get("files") or payload.get("files") or []
|
||||
if not isinstance(files, list):
|
||||
@@ -323,21 +335,24 @@ class CivArchiveClient:
|
||||
return resolved
|
||||
return None
|
||||
|
||||
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]:
|
||||
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||
"""Find model by SHA256 hash value using CivArchive API"""
|
||||
try:
|
||||
payload, error = await self._request_json(f"/sha256/{model_hash.lower()}")
|
||||
if error:
|
||||
if "not found" in error.lower():
|
||||
# Treat a missing payload as an error even when the error string is
|
||||
# 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, error
|
||||
return None, error or "Request failed"
|
||||
|
||||
context, version_data, fallback_files = self._split_context(payload)
|
||||
context, version_data, fallback_files = self._split_context(cast(Dict[str, Any], payload))
|
||||
transformed = self._transform_version(context, version_data, fallback_files)
|
||||
if transformed:
|
||||
return transformed, None
|
||||
|
||||
resolved = await self._resolve_version_from_files(payload)
|
||||
resolved = await self._resolve_version_from_files(cast(Dict[str, Any], payload))
|
||||
if resolved:
|
||||
return resolved, None
|
||||
|
||||
@@ -347,16 +362,30 @@ class CivArchiveClient:
|
||||
except RateLimitError:
|
||||
raise
|
||||
except Exception as 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)
|
||||
|
||||
async def get_model_versions(self, model_id: str) -> Optional[Dict]:
|
||||
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
|
||||
"""Get all versions of a model using CivArchive API"""
|
||||
try:
|
||||
payload, error = await self._request_json(f"/models/{model_id}")
|
||||
if error or payload is None:
|
||||
if error and "not found" in error.lower():
|
||||
return None
|
||||
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
|
||||
|
||||
@@ -364,7 +393,7 @@ class CivArchiveClient:
|
||||
context, version_data, fallback_files = self._split_context(payload)
|
||||
|
||||
versions_meta = data.get("versions") or []
|
||||
transformed_versions: List[Dict] = []
|
||||
transformed_versions: List[Dict[str, Any]] = []
|
||||
for meta in versions_meta:
|
||||
if not isinstance(meta, dict):
|
||||
continue
|
||||
@@ -381,7 +410,7 @@ class CivArchiveClient:
|
||||
if primary_version:
|
||||
transformed_versions.insert(0, primary_version)
|
||||
|
||||
ordered_versions: List[Dict] = []
|
||||
ordered_versions: List[Dict[str, Any]] = []
|
||||
seen_ids = set()
|
||||
for version in transformed_versions:
|
||||
version_id = version.get("id")
|
||||
@@ -402,7 +431,7 @@ class CivArchiveClient:
|
||||
logger.error(f"Error fetching CivArchive model versions for {model_id}: {e}")
|
||||
return None
|
||||
|
||||
async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]:
|
||||
async def get_model_version(self, model_id: int | str | None = None, version_id: int | str | None = None) -> Optional[Dict[str, Any]]:
|
||||
"""Get specific model version using CivArchive API
|
||||
|
||||
Args:
|
||||
@@ -421,6 +450,18 @@ class CivArchiveClient:
|
||||
if error or payload is None:
|
||||
if error and "not found" in error.lower():
|
||||
return None
|
||||
# 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
|
||||
|
||||
@@ -459,7 +500,7 @@ class CivArchiveClient:
|
||||
logger.error(f"Error fetching CivArchive model version via API {model_id}/{version_id}: {e}")
|
||||
return None
|
||||
|
||||
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
|
||||
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||
""" Fetch model version metadata using a known bogus model lookup
|
||||
CivArchive lacks a direct version lookup API, this uses a workaround (which we handle in the main model request now)
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user