mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-14 09:43:22 -03:00
Compare commits
141 Commits
00228deaaa
..
v1.2.0
| Author | SHA1 | Date | |
|---|---|---|---|
| 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 | |||
| bf6a614e0d | |||
| feab01cd9c | |||
| 966024e534 | |||
| 2018722cc8 | |||
| 9d85c2a44a | |||
| 03dd047e62 | |||
| 86b547c1e0 | |||
| bab9752c8b | |||
| 774cc1be86 | |||
| 234b73c8a2 | |||
| abd06c48f4 | |||
| 6ca411e4e4 | |||
| 6470021e77 | |||
| 71658ab37b | |||
| 4f016a8024 | |||
| f362ed585b | |||
| 196172624f | |||
| 316702b7ab | |||
| a7625b009f | |||
| 5d4a33c90d | |||
| 041a6b8525 | |||
| 2638109ad6 | |||
| b019326747 | |||
| 54b44131b6 | |||
| a1d948025c | |||
| a90b2514ba | |||
| cb4ad27813 | |||
| 637831248b | |||
| 2373edf73c | |||
| e0e1b804a7 | |||
| fecbe8241f | |||
| 5983eaa1ce | |||
| 07fa454f72 | |||
| 4b5aa45379 | |||
| 9a0d866be4 | |||
| 308d8f71b8 | |||
| d0e8938039 | |||
| 13ed898b6b | |||
| e1dfd1c2a6 | |||
| e3e944911b | |||
| 51c0135250 | |||
| 7b19bbb14e | |||
| 5494a70f40 | |||
| 26c9ade1c9 | |||
| 87db23825f | |||
| 8fb00998a7 | |||
| dd3aa97d0a | |||
| 905c37290f | |||
| f7632a47f9 | |||
| 646f1ddfb1 | |||
| 170c8068c5 | |||
| a1fd4e150b | |||
| b22f09bd1d | |||
| 4ed9169646 | |||
| f06c60bd47 | |||
| ee8250c26c | |||
| 88349bf944 | |||
| a8adcaf023 | |||
| 63785f82b5 | |||
| cf898da193 |
@@ -36,3 +36,7 @@ vue-widgets/dist/
|
|||||||
|
|
||||||
# Working/research notes (not committed)
|
# Working/research notes (not committed)
|
||||||
.docs/
|
.docs/
|
||||||
|
|
||||||
|
# HF enrichment validation baseline snapshots (contain potentially
|
||||||
|
# NSFW README content fetched from community model repos)
|
||||||
|
tests/enrich_hf_validation/baselines/
|
||||||
|
|||||||
@@ -102,6 +102,7 @@ npm run test:coverage # Generate coverage report
|
|||||||
- ComfyUI: `app.registerExtension()`, `node.addDOMWidget(name, type, element, options)`
|
- ComfyUI: `app.registerExtension()`, `node.addDOMWidget(name, type, element, options)`
|
||||||
- Event handlers via `addEventListener` or widget callbacks
|
- Event handlers via `addEventListener` or widget callbacks
|
||||||
- Shared utilities: `web/comfyui/utils.js`
|
- Shared utilities: `web/comfyui/utils.js`
|
||||||
|
- Dual-mode rendering patterns (canvas vs Vue): see `docs/comfyui-dual-mode-widgets.md`
|
||||||
|
|
||||||
### Vue Composables Pattern
|
### Vue Composables Pattern
|
||||||
|
|
||||||
@@ -136,7 +137,13 @@ npm run test:coverage # Generate coverage report
|
|||||||
- Dual mode: ComfyUI plugin (folder_paths) vs standalone (settings.json)
|
- Dual mode: ComfyUI plugin (folder_paths) vs standalone (settings.json)
|
||||||
- Detection: `os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"`
|
- Detection: `os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"`
|
||||||
- Run `python scripts/sync_translation_keys.py` after adding UI strings to `locales/en.json`
|
- Run `python scripts/sync_translation_keys.py` after adding UI strings to `locales/en.json`
|
||||||
- Symlinks require normalized paths
|
- Symlinks require normalized paths.
|
||||||
|
**Business paths vs real paths**: All stored paths and operation routing use the
|
||||||
|
original paths as they appear under configured model roots — symlinks are NOT
|
||||||
|
resolved. `os.path.realpath` is only for scanner dedup and the symlink cache.
|
||||||
|
Any path passed to `os.remove`/`os.rename`/`shutil.move` or validated by a
|
||||||
|
containment check MUST use the business path (i.e. `os.path.abspath`, not
|
||||||
|
`realpath`).
|
||||||
|
|
||||||
## Git / Commit Messages
|
## Git / Commit Messages
|
||||||
|
|
||||||
|
|||||||
+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_pool import LoraPoolLM
|
||||||
from .py.nodes.lora_randomizer import LoraRandomizerLM
|
from .py.nodes.lora_randomizer import LoraRandomizerLM
|
||||||
from .py.nodes.lora_cycler import LoraCyclerLM
|
from .py.nodes.lora_cycler import LoraCyclerLM
|
||||||
|
from .py.nodes.lora_info import LoraInfoLM
|
||||||
|
from .py.nodes.lora_syntax_to_path import LoraSyntaxToPath
|
||||||
|
from .py.nodes.create_hook_lora import CreateHookLoraLM
|
||||||
|
from .py.nodes.metadata_overwrite import MetadataOverwriteLM
|
||||||
from .py.metadata_collector import init as init_metadata_collector
|
from .py.metadata_collector import init as init_metadata_collector
|
||||||
except (
|
except (
|
||||||
ImportError
|
ImportError
|
||||||
@@ -56,6 +60,16 @@ except (
|
|||||||
"py.nodes.lora_randomizer"
|
"py.nodes.lora_randomizer"
|
||||||
).LoraRandomizerLM
|
).LoraRandomizerLM
|
||||||
LoraCyclerLM = importlib.import_module("py.nodes.lora_cycler").LoraCyclerLM
|
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
|
init_metadata_collector = importlib.import_module("py.metadata_collector").init
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
@@ -75,6 +89,10 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
LoraPoolLM.NAME: LoraPoolLM,
|
LoraPoolLM.NAME: LoraPoolLM,
|
||||||
LoraRandomizerLM.NAME: LoraRandomizerLM,
|
LoraRandomizerLM.NAME: LoraRandomizerLM,
|
||||||
LoraCyclerLM.NAME: LoraCyclerLM,
|
LoraCyclerLM.NAME: LoraCyclerLM,
|
||||||
|
LoraInfoLM.NAME: LoraInfoLM,
|
||||||
|
LoraSyntaxToPath.NAME: LoraSyntaxToPath,
|
||||||
|
CreateHookLoraLM.NAME: CreateHookLoraLM,
|
||||||
|
MetadataOverwriteLM.NAME: MetadataOverwriteLM,
|
||||||
}
|
}
|
||||||
|
|
||||||
WEB_DIRECTORY = "./web/comfyui"
|
WEB_DIRECTORY = "./web/comfyui"
|
||||||
|
|||||||
+346
-295
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,208 @@
|
|||||||
|
# Agent Skills System
|
||||||
|
|
||||||
|
The LoRA Manager agent skills system enables LLM-powered metadata enrichment and other AI-driven tasks. Users configure their own LLM provider (BYOK), and skills are executed through right-click context menu actions.
|
||||||
|
|
||||||
|
## Architecture
|
||||||
|
|
||||||
|
```
|
||||||
|
┌──────────────────────────────────────────────┐
|
||||||
|
│ LoRA Manager Backend │
|
||||||
|
│ │
|
||||||
|
│ ┌──────────────┐ ┌────────────────┐ │
|
||||||
|
│ │ LLMService │───▶│ LLM Provider │ │
|
||||||
|
│ │ (BYOK config, │◀───│ (OpenAI/Ollama │ │
|
||||||
|
│ │ API calls) │ │ /custom) │ │
|
||||||
|
│ └───────┬───────┘ └────────────────┘ │
|
||||||
|
│ │ │
|
||||||
|
│ ┌───────▼───────────────────────┐ │
|
||||||
|
│ │ AgentService │ │
|
||||||
|
│ │ (orchestration: validate │ │
|
||||||
|
│ │ → LLM call → post-process │ │
|
||||||
|
│ │ → WebSocket broadcast) │ │
|
||||||
|
│ └───────┬───────────────────────┘ │
|
||||||
|
│ │ │
|
||||||
|
│ ┌───────▼───────────────────────┐ │
|
||||||
|
│ │ SkillRegistry │ │
|
||||||
|
│ │ ┌─────────────────────────┐ │ │
|
||||||
|
│ │ │ enrich_hf_metadata: │ │ │
|
||||||
|
│ │ │ - skill.yaml │ │ │
|
||||||
|
│ │ │ - prompt.md │ │ │
|
||||||
|
│ │ │ - handler.py │ │ │
|
||||||
|
│ │ └─────────────────────────┘ │ │
|
||||||
|
│ └───────────────────────────────┘ │
|
||||||
|
└──────────────────────────────────────────────┘
|
||||||
|
```
|
||||||
|
|
||||||
|
### Key Design Principle
|
||||||
|
|
||||||
|
**Skills define *what* to do (prompt + post-processing). The AgentService handles *how* (LLM calls, validation, progress).**
|
||||||
|
|
||||||
|
Skills never call the LLM directly. This keeps BYOK configuration centralized and provider-agnostic.
|
||||||
|
|
||||||
|
## BYOK Configuration
|
||||||
|
|
||||||
|
Users configure their LLM provider in **Settings → AI Provider**:
|
||||||
|
|
||||||
|
| Setting | Description | Example |
|
||||||
|
|---|---|---|
|
||||||
|
| `llm_provider` | Provider type | `openai`, `ollama`, or `custom` |
|
||||||
|
| `llm_api_key` | API key (not needed for local Ollama) | `sk-...` |
|
||||||
|
| `llm_api_base` | Custom API base URL (empty = provider default) | `https://api.openai.com/v1` |
|
||||||
|
| `llm_model` | Model name | `gpt-4o-mini` |
|
||||||
|
|
||||||
|
Environment variable overrides: `LLM_API_KEY`, `LLM_MODEL`, `LLM_API_BASE`, `LLM_PROVIDER`.
|
||||||
|
|
||||||
|
### Supported Providers
|
||||||
|
|
||||||
|
- **OpenAI**: Uses `https://api.openai.com/v1` by default
|
||||||
|
- **Ollama** (local): Uses `http://localhost:11434/v1`, no API key required
|
||||||
|
- **Custom**: Any OpenAI-compatible endpoint (vLLM, LM Studio, etc.) — set `llm_api_base` explicitly
|
||||||
|
|
||||||
|
## Available Skills
|
||||||
|
|
||||||
|
### enrich_hf_metadata
|
||||||
|
|
||||||
|
Enriches HuggingFace-downloaded models with metadata extracted by an LLM from the HF model card.
|
||||||
|
|
||||||
|
**Entry point**: Right-click context menu → "Enrich Metadata (Agent)"
|
||||||
|
|
||||||
|
**What it does**:
|
||||||
|
1. Reads the model's `.metadata.json` to get the `hf_url`
|
||||||
|
2. Fetches the README.md from the HuggingFace repository
|
||||||
|
3. Sends the README + local metadata to the LLM for structured extraction
|
||||||
|
4. Writes extracted fields to `.metadata.json`:
|
||||||
|
- `base_model` — only if current value is empty
|
||||||
|
- `trainedWords` — trigger words (LoRA only, if none exist)
|
||||||
|
- `modelDescription` — concise summary (if none exists)
|
||||||
|
- `tags` — merged with existing tags, deduplicated
|
||||||
|
- `metadata_source` — audit trail: `agent:enrich_hf_metadata`
|
||||||
|
- `llm_enriched_at` — ISO timestamp
|
||||||
|
5. Downloads and optimizes preview image (if LLM found one in the README)
|
||||||
|
6. Updates the scanner cache
|
||||||
|
7. Broadcasts WebSocket progress events
|
||||||
|
|
||||||
|
**Model types**: LoRA, Checkpoint, Embedding
|
||||||
|
|
||||||
|
## Adding a New Skill
|
||||||
|
|
||||||
|
### 1. Create the skill directory
|
||||||
|
|
||||||
|
```
|
||||||
|
py/services/agent/skills/<skill_name>/
|
||||||
|
├── skill.yaml # Skill metadata and schemas
|
||||||
|
├── prompt.md # LLM prompt template
|
||||||
|
└── handler.py # Pre-processing and post-processing
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2. Write skill.yaml
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
name: my_skill
|
||||||
|
title: "My Skill"
|
||||||
|
description: "What this skill does"
|
||||||
|
llm_required: true
|
||||||
|
model_type_filter: ["lora"] # or null for all types
|
||||||
|
input_schema:
|
||||||
|
type: object
|
||||||
|
properties:
|
||||||
|
model_paths:
|
||||||
|
type: array
|
||||||
|
items:
|
||||||
|
type: string
|
||||||
|
required:
|
||||||
|
- model_paths
|
||||||
|
output_schema:
|
||||||
|
type: object
|
||||||
|
properties:
|
||||||
|
# ... JSON schema for LLM output
|
||||||
|
permissions:
|
||||||
|
write_metadata: true
|
||||||
|
write_previews: false
|
||||||
|
network_domains:
|
||||||
|
- "example.com"
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3. Write prompt.md
|
||||||
|
|
||||||
|
Use `{{variable}}` placeholders that will be replaced with data from the `prepare` function:
|
||||||
|
|
||||||
|
```markdown
|
||||||
|
You are an expert assistant...
|
||||||
|
|
||||||
|
Model URL: {{hf_url}}
|
||||||
|
README content:
|
||||||
|
{{readme_content}}
|
||||||
|
|
||||||
|
Current metadata:
|
||||||
|
{{current_metadata}}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 4. Write handler.py
|
||||||
|
|
||||||
|
```python
|
||||||
|
async def prepare(model_path: str, input_data: dict) -> dict:
|
||||||
|
"""Gather context for the LLM prompt. Returns variables for template rendering."""
|
||||||
|
return {
|
||||||
|
"model_path": model_path,
|
||||||
|
# ... other variables used in prompt.md
|
||||||
|
}
|
||||||
|
|
||||||
|
async def post_process(context) -> dict:
|
||||||
|
"""Apply the LLM-extracted data to the model."""
|
||||||
|
llm_response = context.llm_response
|
||||||
|
# ... write metadata, download previews, update cache
|
||||||
|
return {
|
||||||
|
"success": True,
|
||||||
|
"updated_fields": ["base_model", "tags"],
|
||||||
|
"errors": [],
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Important**: Use absolute imports (`from py.utils.metadata_manager import MetadataManager`) because skills are loaded via `importlib.util.spec_from_file_location`, which doesn't support relative imports.
|
||||||
|
|
||||||
|
### 5. Test
|
||||||
|
|
||||||
|
The skill is automatically discovered by `SkillRegistry` on startup. Test with:
|
||||||
|
|
||||||
|
```python
|
||||||
|
pytest tests/services/test_agent_service.py
|
||||||
|
```
|
||||||
|
|
||||||
|
## API Endpoints
|
||||||
|
|
||||||
|
| Method | Path | Description |
|
||||||
|
|---|---|---|
|
||||||
|
| GET | `/api/lm/agent/skills` | List available skills |
|
||||||
|
| POST | `/api/lm/agent/execute/{skill_name}` | Execute a skill (body: `{"model_paths": [...]}`) |
|
||||||
|
| POST | `/api/lm/agent/cancel` | Cancel running skill (stub) |
|
||||||
|
|
||||||
|
## WebSocket Events
|
||||||
|
|
||||||
|
| Type | When | Key fields |
|
||||||
|
|---|---|---|
|
||||||
|
| `agent_progress` | Skill started/processing | `skill`, `status`, `total`, `processed`, `success`, `current_path` |
|
||||||
|
| `agent_progress` | Skill completed | `skill`, `status`, `updated_models`, `errors`, `summary` |
|
||||||
|
| `agent_progress` | Skill error | `skill`, `status`, `error` |
|
||||||
|
|
||||||
|
## Security Model
|
||||||
|
|
||||||
|
Skills declare permissions in `skill.yaml`:
|
||||||
|
- `write_metadata` — can write `.metadata.json` files
|
||||||
|
- `write_previews` — can download/replace preview images
|
||||||
|
- `network_domains` — allowed domains for HTTP requests
|
||||||
|
|
||||||
|
These are declarative constraints checked by `AgentService`. They are defense-in-depth, not a sandbox — the Python process can technically do anything, but the contract is clear and auditable.
|
||||||
|
|
||||||
|
## File Locations
|
||||||
|
|
||||||
|
| Component | Path |
|
||||||
|
|---|---|
|
||||||
|
| LLMService | `py/services/llm_service.py` |
|
||||||
|
| AgentService | `py/services/agent/agent_service.py` |
|
||||||
|
| SkillRegistry | `py/services/agent/skill_registry.py` |
|
||||||
|
| SkillDefinition | `py/services/agent/skill_definition.py` |
|
||||||
|
| Skills directory | `py/services/agent/skills/` |
|
||||||
|
| Route handlers | `py/routes/handlers/agent_handlers.py` |
|
||||||
|
| Frontend manager | `static/js/managers/AgentManager.js` |
|
||||||
|
| Settings UI | `templates/components/modals/settings_modal.html` |
|
||||||
|
| Context menu | `templates/components/context_menu.html` |
|
||||||
@@ -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
|
||||||
File diff suppressed because one or more lines are too long
+2218
-2143
File diff suppressed because it is too large
Load Diff
+2218
-2143
File diff suppressed because it is too large
Load Diff
+2218
-2143
File diff suppressed because it is too large
Load Diff
+2218
-2143
File diff suppressed because it is too large
Load Diff
+2218
-2143
File diff suppressed because it is too large
Load Diff
+2218
-2143
File diff suppressed because it is too large
Load Diff
+2218
-2143
File diff suppressed because it is too large
Load Diff
+2218
-2143
File diff suppressed because it is too large
Load Diff
+2218
-2143
File diff suppressed because it is too large
Load Diff
+2218
-2143
File diff suppressed because it is too large
Load Diff
+68
-8
@@ -8,6 +8,8 @@ from typing import Any, Dict, Iterable, List, Mapping, Optional, Set, Tuple
|
|||||||
import logging
|
import logging
|
||||||
import json
|
import json
|
||||||
import urllib.parse
|
import urllib.parse
|
||||||
|
import sys as _sys
|
||||||
|
import types as _types
|
||||||
import time
|
import time
|
||||||
|
|
||||||
from .utils.cache_paths import CacheType, get_cache_file_path, get_legacy_cache_paths
|
from .utils.cache_paths import CacheType, get_cache_file_path, get_legacy_cache_paths
|
||||||
@@ -175,8 +177,7 @@ class Config:
|
|||||||
|
|
||||||
# Load extra folder paths from active library settings before symlink scan
|
# Load extra folder paths from active library settings before symlink scan
|
||||||
# so both primary and extra paths are discovered in a single pass.
|
# so both primary and extra paths are discovered in a single pass.
|
||||||
if not standalone_mode:
|
self._load_extra_paths_from_settings()
|
||||||
self._load_extra_paths_from_settings()
|
|
||||||
|
|
||||||
# Scan symbolic links during initialization
|
# Scan symbolic links during initialization
|
||||||
self._initialize_symlink_mappings()
|
self._initialize_symlink_mappings()
|
||||||
@@ -191,7 +192,7 @@ class Config:
|
|||||||
Called during ``Config.__init__`` before the symlink scan so both primary and
|
Called during ``Config.__init__`` before the symlink scan so both primary and
|
||||||
extra paths are discovered in a single pass. Mirrors the extra-path
|
extra paths are discovered in a single pass. Mirrors the extra-path
|
||||||
portion of ``_apply_library_paths`` without replacing the primary roots
|
portion of ``_apply_library_paths`` without replacing the primary roots
|
||||||
that were already resolved from ComfyUI's ``folder_paths``.
|
that were already resolved via ``folder_paths.get_folder_paths``.
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
from .services.settings_manager import get_settings_manager
|
from .services.settings_manager import get_settings_manager
|
||||||
@@ -207,6 +208,12 @@ class Config:
|
|||||||
if not isinstance(library_config, dict):
|
if not isinstance(library_config, dict):
|
||||||
return
|
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")
|
extra_folder_paths = library_config.get("extra_folder_paths")
|
||||||
if not isinstance(extra_folder_paths, dict):
|
if not isinstance(extra_folder_paths, dict):
|
||||||
return
|
return
|
||||||
@@ -232,10 +239,6 @@ class Config:
|
|||||||
extra_embedding
|
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:
|
if self.extra_loras_roots:
|
||||||
logger.info(
|
logger.info(
|
||||||
"Found extra LoRA roots:"
|
"Found extra LoRA roots:"
|
||||||
@@ -356,6 +359,47 @@ class Config:
|
|||||||
"Failed to rename legacy 'default' library: %s", rename_error
|
"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(
|
default_lora_root = _resolve_valid_default_root(
|
||||||
comfy_library.get("default_lora_root", ""),
|
comfy_library.get("default_lora_root", ""),
|
||||||
list(self.loras_roots or []),
|
list(self.loras_roots or []),
|
||||||
@@ -1380,4 +1424,20 @@ class Config:
|
|||||||
|
|
||||||
|
|
||||||
# Global config instance
|
# Global config instance
|
||||||
config = Config()
|
# NOTE: Guard against re-import. When ServiceRegistry.get_lora_scanner() triggers
|
||||||
|
# a fresh import of lora_scanner → config, we must NOT re-execute Config.__init__()
|
||||||
|
# (which re-scans all roots, re-registers libraries, etc.).
|
||||||
|
#
|
||||||
|
# Strategy: store the config instance in a dedicated sentinel module
|
||||||
|
# ('_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"
|
||||||
|
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]
|
||||||
|
else:
|
||||||
|
config: Config = Config()
|
||||||
|
# Register the sentinel so re-imports of py.config find us.
|
||||||
|
_sentinel_mod = _types.ModuleType(_CONFIG_SENTINEL)
|
||||||
|
_sentinel_mod.config = config
|
||||||
|
_sys.modules[_CONFIG_SENTINEL] = _sentinel_mod
|
||||||
|
|||||||
@@ -208,6 +208,10 @@ class LoraManager:
|
|||||||
# Initialize WebSocket manager
|
# Initialize WebSocket manager
|
||||||
await ServiceRegistry.get_websocket_manager()
|
await ServiceRegistry.get_websocket_manager()
|
||||||
|
|
||||||
|
# Preload LLM model catalog (background task, non-blocking)
|
||||||
|
from .services.llm_service import LLMService
|
||||||
|
await LLMService.get_instance()
|
||||||
|
|
||||||
# Initialize scanners in background
|
# Initialize scanners in background
|
||||||
lora_scanner = await ServiceRegistry.get_lora_scanner()
|
lora_scanner = await ServiceRegistry.get_lora_scanner()
|
||||||
checkpoint_scanner = await ServiceRegistry.get_checkpoint_scanner()
|
checkpoint_scanner = await ServiceRegistry.get_checkpoint_scanner()
|
||||||
@@ -445,5 +449,12 @@ class LoraManager:
|
|||||||
scanner.cancel_task()
|
scanner.cancel_task()
|
||||||
logger.debug("LoRA Manager: Cancelled %s", name)
|
logger.debug("LoRA Manager: Cancelled %s", name)
|
||||||
|
|
||||||
|
# Close shared aiohttp sessions to avoid "Unclosed client session" warnings
|
||||||
|
try:
|
||||||
|
from py.routes.handlers.hf_handlers import close_hf_api_session
|
||||||
|
await close_hf_api_session()
|
||||||
|
except Exception as exc:
|
||||||
|
logger.debug("Error closing HF API session: %s", exc)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error during cleanup: {e}", exc_info=True)
|
logger.error(f"Error during cleanup: {e}", exc_info=True)
|
||||||
|
|||||||
@@ -1,5 +1,11 @@
|
|||||||
"""Constants used by the metadata collector"""
|
"""Constants used by the metadata collector"""
|
||||||
|
|
||||||
|
# Sentinel value for clip_skip to distinguish "unconnected / widget default"
|
||||||
|
# from "user wired value 0". Both ComfyUI CLIPSetLastLayer (-24..-1) and
|
||||||
|
# A1111 conventions treat 0 as meaningless for clip skipping, but users may
|
||||||
|
# explicitly wire 0 to the overwrite node to express "no clip skip / default".
|
||||||
|
CLIP_SKIP_SENTINEL = -25
|
||||||
|
|
||||||
# Metadata categories
|
# Metadata categories
|
||||||
MODELS = "models"
|
MODELS = "models"
|
||||||
PROMPTS = "prompts"
|
PROMPTS = "prompts"
|
||||||
@@ -9,6 +15,14 @@ EMBEDDINGS = "embeddings"
|
|||||||
SIZE = "size"
|
SIZE = "size"
|
||||||
IMAGES = "images"
|
IMAGES = "images"
|
||||||
IS_SAMPLER = "is_sampler" # New constant to mark sampler nodes
|
IS_SAMPLER = "is_sampler" # New constant to mark sampler nodes
|
||||||
|
OVERWRITE = "overwrite" # Manual metadata overwrite from MetadataOverwriteLM node
|
||||||
|
|
||||||
|
# Field names that the MetadataOverwriteLM node and its extractor share
|
||||||
|
METADATA_OVERWRITE_FIELDS = (
|
||||||
|
"prompt", "negative_prompt", "seed", "steps", "cfg_scale",
|
||||||
|
"sampler", "scheduler", "model", "loras", "size",
|
||||||
|
"clip_skip", "additional_data",
|
||||||
|
)
|
||||||
|
|
||||||
# Complete list of categories to track
|
# Complete list of categories to track
|
||||||
METADATA_CATEGORIES = [MODELS, PROMPTS, SAMPLING, LORAS, EMBEDDINGS, SIZE, IMAGES]
|
METADATA_CATEGORIES = [MODELS, PROMPTS, SAMPLING, LORAS, EMBEDDINGS, SIZE, IMAGES, OVERWRITE]
|
||||||
|
|||||||
@@ -83,7 +83,8 @@ class MetadataHook:
|
|||||||
|
|
||||||
# Record inputs before execution
|
# Record inputs before execution
|
||||||
if node_id is not None:
|
if node_id is not None:
|
||||||
registry.record_node_execution(node_id, class_type, input_data_all, None)
|
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||||
|
registry.record_node_execution(node_id, class_type, input_data_all, None, return_types=return_types)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
||||||
|
|
||||||
@@ -114,7 +115,8 @@ class MetadataHook:
|
|||||||
|
|
||||||
# Record outputs after execution
|
# Record outputs after execution
|
||||||
if node_id is not None:
|
if node_id is not None:
|
||||||
registry.update_node_execution(node_id, class_type, results)
|
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||||
|
registry.update_node_execution(node_id, class_type, results, return_types=return_types)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
||||||
|
|
||||||
@@ -136,6 +138,9 @@ class MetadataHook:
|
|||||||
if hasattr(prompt, 'original_prompt'):
|
if hasattr(prompt, 'original_prompt'):
|
||||||
registry.set_current_prompt(prompt)
|
registry.set_current_prompt(prompt)
|
||||||
|
|
||||||
|
# Store extra_data for accessing full workflow node properties
|
||||||
|
registry.set_extra_data(extra_data)
|
||||||
|
|
||||||
# Execute the original function
|
# Execute the original function
|
||||||
return original_execute(*args, **kwargs)
|
return original_execute(*args, **kwargs)
|
||||||
|
|
||||||
@@ -163,7 +168,8 @@ class MetadataHook:
|
|||||||
class_type = obj.__class__.__name__
|
class_type = obj.__class__.__name__
|
||||||
node_id = unique_id
|
node_id = unique_id
|
||||||
if node_id is not None:
|
if node_id is not None:
|
||||||
registry.record_node_execution(node_id, class_type, input_data_all, None)
|
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||||
|
registry.record_node_execution(node_id, class_type, input_data_all, None, return_types=return_types)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
||||||
|
|
||||||
@@ -180,7 +186,8 @@ class MetadataHook:
|
|||||||
class_type = obj.__class__.__name__
|
class_type = obj.__class__.__name__
|
||||||
node_id = unique_id
|
node_id = unique_id
|
||||||
if node_id is not None:
|
if node_id is not None:
|
||||||
registry.update_node_execution(node_id, class_type, results)
|
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||||
|
registry.update_node_execution(node_id, class_type, results, return_types=return_types)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
||||||
|
|
||||||
@@ -202,6 +209,9 @@ class MetadataHook:
|
|||||||
if hasattr(prompt, 'original_prompt'):
|
if hasattr(prompt, 'original_prompt'):
|
||||||
registry.set_current_prompt(prompt)
|
registry.set_current_prompt(prompt)
|
||||||
|
|
||||||
|
# Store extra_data for accessing full workflow node properties
|
||||||
|
registry.set_extra_data(extra_data)
|
||||||
|
|
||||||
# Execute the original function
|
# Execute the original function
|
||||||
return await original_execute(*args, **kwargs)
|
return await original_execute(*args, **kwargs)
|
||||||
|
|
||||||
|
|||||||
@@ -1,15 +1,68 @@
|
|||||||
import json
|
import json
|
||||||
|
import logging
|
||||||
import os
|
import os
|
||||||
from .constants import IMAGES
|
from .constants import IMAGES
|
||||||
|
|
||||||
# Check if running in standalone mode
|
# Check if running in standalone mode
|
||||||
standalone_mode = os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1" or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0"
|
standalone_mode = os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1" or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0"
|
||||||
|
|
||||||
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IS_SAMPLER
|
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IS_SAMPLER, OVERWRITE
|
||||||
|
from .node_extractors import NODE_EXTRACTORS
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Keys that identify metadata hint marks stored in node.properties.lm_marker_role
|
||||||
|
_META_MARK_PREFIX = "meta_"
|
||||||
|
_MARK_PRIMARY_MODEL = "primary_model"
|
||||||
|
_MARK_PRIMARY_SAMPLER = "primary_sampler"
|
||||||
|
_MARK_POSITIVE_PROMPT = "positive_prompt"
|
||||||
|
_MARK_NEGATIVE_PROMPT = "negative_prompt"
|
||||||
|
|
||||||
class MetadataProcessor:
|
class MetadataProcessor:
|
||||||
"""Process and format collected metadata"""
|
"""Process and format collected metadata"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_user_marks(metadata):
|
||||||
|
"""Scan workflow nodes (from extra_data.extra_pnginfo.workflow) for user-assigned
|
||||||
|
metadata hint marks stored in node.properties.lm_marker_role.
|
||||||
|
|
||||||
|
Returns a dict mapping mark type keys to node IDs.
|
||||||
|
Example: {'primary_model': '42', 'primary_sampler': '17'}
|
||||||
|
"""
|
||||||
|
marks: dict[str, str] = {}
|
||||||
|
|
||||||
|
# Primary source: extra_data.extra_pnginfo.workflow.nodes (has full properties)
|
||||||
|
extra_data = metadata.get("extra_data")
|
||||||
|
if extra_data and isinstance(extra_data, dict):
|
||||||
|
extra_pnginfo = extra_data.get("extra_pnginfo", {})
|
||||||
|
if isinstance(extra_pnginfo, dict):
|
||||||
|
workflow = extra_pnginfo.get("workflow", {})
|
||||||
|
nodes = workflow.get("nodes", [])
|
||||||
|
for node in nodes:
|
||||||
|
node_id = str(node.get("id", ""))
|
||||||
|
role = node.get("properties", {}).get("lm_marker_role", "")
|
||||||
|
if role.startswith(_META_MARK_PREFIX):
|
||||||
|
mark_type = role[len(_META_MARK_PREFIX):]
|
||||||
|
if mark_type in marks:
|
||||||
|
logger.warning(
|
||||||
|
"Duplicate meta hint '%s': node %s (previous: %s), "
|
||||||
|
"last match wins",
|
||||||
|
mark_type, node_id, marks[mark_type],
|
||||||
|
)
|
||||||
|
marks[mark_type] = node_id
|
||||||
|
|
||||||
|
# Fallback: try prompt.original_prompt (API-only submissions may not have workflow)
|
||||||
|
if not marks:
|
||||||
|
prompt = metadata.get("current_prompt")
|
||||||
|
if prompt and getattr(prompt, "original_prompt", None):
|
||||||
|
for node_id, node_data in prompt.original_prompt.items():
|
||||||
|
role = node_data.get("properties", {}).get("lm_marker_role", "")
|
||||||
|
if role.startswith(_META_MARK_PREFIX):
|
||||||
|
mark_type = role[len(_META_MARK_PREFIX):]
|
||||||
|
marks[mark_type] = node_id
|
||||||
|
|
||||||
|
return marks
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def find_primary_sampler(metadata, downstream_id=None):
|
def find_primary_sampler(metadata, downstream_id=None):
|
||||||
"""
|
"""
|
||||||
@@ -471,20 +524,57 @@ class MetadataProcessor:
|
|||||||
"checkpoint": None,
|
"checkpoint": None,
|
||||||
"loras": "",
|
"loras": "",
|
||||||
"size": None,
|
"size": None,
|
||||||
"clip_skip": None
|
"clip_skip": None,
|
||||||
|
"additional_data": "",
|
||||||
}
|
}
|
||||||
|
|
||||||
# Get the prompt object for node relationship tracing
|
# Get the prompt object for node relationship tracing
|
||||||
prompt = metadata.get("current_prompt")
|
prompt = metadata.get("current_prompt")
|
||||||
|
|
||||||
# Find the primary KSampler node
|
# ---- User marks: override heuristic inference with user-assigned hints ----
|
||||||
primary_sampler_id, primary_sampler = MetadataProcessor.find_primary_sampler(metadata, id)
|
user_marks = MetadataProcessor._get_user_marks(metadata)
|
||||||
|
|
||||||
# Directly get checkpoint from metadata instead of tracing
|
# Find the primary KSampler node (user mark takes priority)
|
||||||
# Pass primary_sampler_id to avoid redundant calculation
|
primary_sampler_id = None
|
||||||
checkpoint = MetadataProcessor.find_primary_checkpoint(metadata, id, primary_sampler_id)
|
primary_sampler = None
|
||||||
if checkpoint:
|
if _MARK_PRIMARY_SAMPLER in user_marks:
|
||||||
params["checkpoint"] = checkpoint
|
marked_id = user_marks[_MARK_PRIMARY_SAMPLER]
|
||||||
|
sampler_data = metadata.get(SAMPLING, {}).get(marked_id)
|
||||||
|
if sampler_data and sampler_data.get(IS_SAMPLER):
|
||||||
|
primary_sampler_id = marked_id
|
||||||
|
primary_sampler = sampler_data
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"User-marked primary sampler %s has no runtime metadata, "
|
||||||
|
"falling back to heuristic",
|
||||||
|
marked_id,
|
||||||
|
)
|
||||||
|
if primary_sampler is None:
|
||||||
|
primary_sampler_id, primary_sampler = MetadataProcessor.find_primary_sampler(metadata, id)
|
||||||
|
|
||||||
|
# Resolve checkpoint / model (user mark takes priority)
|
||||||
|
if _MARK_PRIMARY_MODEL in user_marks:
|
||||||
|
marked_id = user_marks[_MARK_PRIMARY_MODEL]
|
||||||
|
if marked_id in metadata.get(MODELS, {}):
|
||||||
|
params["checkpoint"] = metadata[MODELS][marked_id].get("name")
|
||||||
|
else:
|
||||||
|
extra_data = metadata.get("extra_data")
|
||||||
|
extra_pnginfo = extra_data.get("extra_pnginfo", {}) if extra_data and isinstance(extra_data, dict) else {}
|
||||||
|
workflow = extra_pnginfo.get("workflow", {}) if isinstance(extra_pnginfo, dict) else {}
|
||||||
|
node_type = "unknown"
|
||||||
|
for n in workflow.get("nodes", []):
|
||||||
|
if str(n.get("id", "")) == marked_id:
|
||||||
|
node_type = n.get("type", "unknown")
|
||||||
|
break
|
||||||
|
logger.warning(
|
||||||
|
"User-marked primary model %s (type=%s, registered=%s) has no runtime metadata, "
|
||||||
|
"falling back to heuristic",
|
||||||
|
marked_id, node_type, node_type in NODE_EXTRACTORS,
|
||||||
|
)
|
||||||
|
if params["checkpoint"] is None:
|
||||||
|
checkpoint = MetadataProcessor.find_primary_checkpoint(metadata, id, primary_sampler_id)
|
||||||
|
if checkpoint:
|
||||||
|
params["checkpoint"] = checkpoint
|
||||||
|
|
||||||
# Check if guidance parameter exists in any sampling node
|
# Check if guidance parameter exists in any sampling node
|
||||||
for node_id, sampler_info in metadata.get(SAMPLING, {}).items():
|
for node_id, sampler_info in metadata.get(SAMPLING, {}).items():
|
||||||
@@ -540,6 +630,21 @@ class MetadataProcessor:
|
|||||||
# For SamplerCustom, handle any additional parameters
|
# For SamplerCustom, handle any additional parameters
|
||||||
MetadataProcessor.handle_custom_advanced_sampler(metadata, prompt, primary_sampler_id, params)
|
MetadataProcessor.handle_custom_advanced_sampler(metadata, prompt, primary_sampler_id, params)
|
||||||
|
|
||||||
|
# ---- User marks: override prompts with explicitly tagged nodes ----
|
||||||
|
prompts_data = metadata.get(PROMPTS, {})
|
||||||
|
if _MARK_POSITIVE_PROMPT in user_marks:
|
||||||
|
pos_id = user_marks[_MARK_POSITIVE_PROMPT]
|
||||||
|
if pos_id in prompts_data:
|
||||||
|
prompt_text = prompts_data[pos_id].get("text") or prompts_data[pos_id].get("positive_text")
|
||||||
|
if prompt_text:
|
||||||
|
params["prompt"] = prompt_text
|
||||||
|
if _MARK_NEGATIVE_PROMPT in user_marks:
|
||||||
|
neg_id = user_marks[_MARK_NEGATIVE_PROMPT]
|
||||||
|
if neg_id in prompts_data:
|
||||||
|
prompt_text = prompts_data[neg_id].get("text") or prompts_data[neg_id].get("negative_text")
|
||||||
|
if prompt_text:
|
||||||
|
params["negative_prompt"] = prompt_text
|
||||||
|
|
||||||
# Size extraction is same for all sampler types
|
# Size extraction is same for all sampler types
|
||||||
# Check if the sampler itself has size information (from latent_image)
|
# Check if the sampler itself has size information (from latent_image)
|
||||||
if primary_sampler_id in metadata.get(SIZE, {}):
|
if primary_sampler_id in metadata.get(SIZE, {}):
|
||||||
@@ -569,6 +674,25 @@ class MetadataProcessor:
|
|||||||
if params["clip_skip"] is None:
|
if params["clip_skip"] is None:
|
||||||
params["clip_skip"] = "1"
|
params["clip_skip"] = "1"
|
||||||
|
|
||||||
|
# ---- Apply manual metadata overwrites ----
|
||||||
|
for overwrite_info in metadata.get(OVERWRITE, {}).values():
|
||||||
|
overwrite_params = overwrite_info.get("parameters", {})
|
||||||
|
for key, value in overwrite_params.items():
|
||||||
|
if key == "clip_skip":
|
||||||
|
# Accept any value from overwrite node (sentinel -25 already
|
||||||
|
# filtered upstream). Needed because falsy check treats 0
|
||||||
|
# as "not set" even though 0 is a valid wired input here.
|
||||||
|
params[key] = value
|
||||||
|
elif value: # truthy check — only overwrite when user provided a real value
|
||||||
|
params[key] = value
|
||||||
|
|
||||||
|
# Bridge: the overwrite node exposes the field as "model" (more accurate),
|
||||||
|
# but the internal pipeline key remains "checkpoint" for backward compatibility
|
||||||
|
# with A1111 metadata format and downstream consumers.
|
||||||
|
if params.get("model"):
|
||||||
|
params["checkpoint"] = params["model"]
|
||||||
|
del params["model"]
|
||||||
|
|
||||||
return params
|
return params
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
import time
|
import time
|
||||||
from nodes import NODE_CLASS_MAPPINGS # type: ignore
|
from nodes import NODE_CLASS_MAPPINGS # type: ignore
|
||||||
from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor
|
from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor
|
||||||
from .constants import METADATA_CATEGORIES, IMAGES
|
from .constants import METADATA_CATEGORIES, IMAGES, OVERWRITE
|
||||||
|
|
||||||
|
|
||||||
class MetadataRegistry:
|
class MetadataRegistry:
|
||||||
@@ -61,6 +61,7 @@ class MetadataRegistry:
|
|||||||
{
|
{
|
||||||
"execution_order": [],
|
"execution_order": [],
|
||||||
"current_prompt": None, # Will store the prompt object
|
"current_prompt": None, # Will store the prompt object
|
||||||
|
"extra_data": None, # Will store the API extra_data for workflow metadata
|
||||||
"timestamp": time.time(),
|
"timestamp": time.time(),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
@@ -75,6 +76,11 @@ class MetadataRegistry:
|
|||||||
# Store the prompt in the metadata for later relationship tracing
|
# Store the prompt in the metadata for later relationship tracing
|
||||||
self.prompt_metadata[self.current_prompt_id]["current_prompt"] = prompt
|
self.prompt_metadata[self.current_prompt_id]["current_prompt"] = prompt
|
||||||
|
|
||||||
|
def set_extra_data(self, extra_data):
|
||||||
|
"""Store the API extra_data (contains extra_pnginfo.workflow with node properties)"""
|
||||||
|
if self.current_prompt_id and self.current_prompt_id in self.prompt_metadata:
|
||||||
|
self.prompt_metadata[self.current_prompt_id]["extra_data"] = extra_data
|
||||||
|
|
||||||
def get_metadata(self, prompt_id=None):
|
def get_metadata(self, prompt_id=None):
|
||||||
"""Get collected metadata for a prompt"""
|
"""Get collected metadata for a prompt"""
|
||||||
key = prompt_id if prompt_id is not None else self.current_prompt_id
|
key = prompt_id if prompt_id is not None else self.current_prompt_id
|
||||||
@@ -122,20 +128,28 @@ class MetadataRegistry:
|
|||||||
cache_key = f"{node_id}:{class_type}"
|
cache_key = f"{node_id}:{class_type}"
|
||||||
|
|
||||||
# Check if this node type is relevant for metadata collection
|
# Check if this node type is relevant for metadata collection
|
||||||
if class_type in NODE_EXTRACTORS:
|
if class_type in NODE_EXTRACTORS or cache_key in self.node_cache:
|
||||||
# Check if we have cached metadata for this node
|
# Check if we have cached metadata for this node
|
||||||
if cache_key in self.node_cache:
|
if cache_key in self.node_cache:
|
||||||
cached_data = self.node_cache[cache_key]
|
cached_data = self.node_cache[cache_key]
|
||||||
|
|
||||||
|
# Detect bypass (mode=4) / mute (mode=2) — these nodes
|
||||||
|
# were intentionally disabled and should not contribute
|
||||||
|
# overwrite values from a previous execution's cache.
|
||||||
|
node_mode = node_data.get("mode", 0)
|
||||||
|
node_is_disabled = node_mode in (2, 4)
|
||||||
|
|
||||||
# Apply cached metadata to the current metadata
|
# Apply cached metadata to the current metadata
|
||||||
for category in self.metadata_categories:
|
for category in self.metadata_categories:
|
||||||
|
if category == OVERWRITE and node_is_disabled:
|
||||||
|
continue
|
||||||
if category in cached_data and node_id in cached_data[category]:
|
if category in cached_data and node_id in cached_data[category]:
|
||||||
if node_id not in metadata[category]:
|
if node_id not in metadata[category]:
|
||||||
metadata[category][node_id] = cached_data[category][
|
metadata[category][node_id] = cached_data[category][
|
||||||
node_id
|
node_id
|
||||||
]
|
]
|
||||||
|
|
||||||
def record_node_execution(self, node_id, class_type, inputs, outputs):
|
def record_node_execution(self, node_id, class_type, inputs, outputs, return_types=None):
|
||||||
"""Record information about a node's execution"""
|
"""Record information about a node's execution"""
|
||||||
if not self.current_prompt_id:
|
if not self.current_prompt_id:
|
||||||
return
|
return
|
||||||
@@ -158,17 +172,18 @@ class MetadataRegistry:
|
|||||||
|
|
||||||
# Extract node-specific metadata
|
# Extract node-specific metadata
|
||||||
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
||||||
extractor.extract(
|
if extractor is GenericNodeExtractor:
|
||||||
node_id,
|
extractor.extract(node_id, processed_inputs, outputs,
|
||||||
processed_inputs,
|
self.prompt_metadata[self.current_prompt_id],
|
||||||
outputs,
|
return_types=return_types)
|
||||||
self.prompt_metadata[self.current_prompt_id],
|
else:
|
||||||
)
|
extractor.extract(node_id, processed_inputs, outputs,
|
||||||
|
self.prompt_metadata[self.current_prompt_id])
|
||||||
|
|
||||||
# Cache this node's metadata
|
# Cache this node's metadata
|
||||||
self._cache_node_metadata(node_id, class_type)
|
self._cache_node_metadata(node_id, class_type)
|
||||||
|
|
||||||
def update_node_execution(self, node_id, class_type, outputs):
|
def update_node_execution(self, node_id, class_type, outputs, return_types=None):
|
||||||
"""Update node metadata with output information"""
|
"""Update node metadata with output information"""
|
||||||
if not self.current_prompt_id:
|
if not self.current_prompt_id:
|
||||||
return
|
return
|
||||||
@@ -179,9 +194,17 @@ class MetadataRegistry:
|
|||||||
# Use the same extractor to update with outputs
|
# Use the same extractor to update with outputs
|
||||||
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
||||||
if hasattr(extractor, "update"):
|
if hasattr(extractor, "update"):
|
||||||
extractor.update(
|
if extractor is GenericNodeExtractor:
|
||||||
node_id, processed_outputs, self.prompt_metadata[self.current_prompt_id]
|
extractor.update(
|
||||||
)
|
node_id, processed_outputs,
|
||||||
|
self.prompt_metadata[self.current_prompt_id],
|
||||||
|
return_types=return_types,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
extractor.update(
|
||||||
|
node_id, processed_outputs,
|
||||||
|
self.prompt_metadata[self.current_prompt_id],
|
||||||
|
)
|
||||||
|
|
||||||
# Update the cached metadata for this node
|
# Update the cached metadata for this node
|
||||||
self._cache_node_metadata(node_id, class_type)
|
self._cache_node_metadata(node_id, class_type)
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ import json
|
|||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
|
|
||||||
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER
|
from .constants import CLIP_SKIP_SENTINEL, MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER, OVERWRITE, METADATA_OVERWRITE_FIELDS
|
||||||
|
|
||||||
|
|
||||||
def _store_checkpoint_metadata(metadata, node_id, model_name):
|
def _store_checkpoint_metadata(metadata, node_id, model_name):
|
||||||
@@ -31,10 +31,77 @@ class NodeMetadataExtractor:
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
class GenericNodeExtractor(NodeMetadataExtractor):
|
class GenericNodeExtractor(NodeMetadataExtractor):
|
||||||
"""Default extractor for nodes without specific handling"""
|
"""Fallback extractor with type-signature-based detection.
|
||||||
|
|
||||||
|
When a node is not in the NODE_EXTRACTORS registry, the hook layer
|
||||||
|
passes ``return_types`` from ``obj.RETURN_TYPES``:
|
||||||
|
|
||||||
|
* ``MODEL`` output: common input fields (ckpt_name, unet_name, etc.)
|
||||||
|
are checked for a model file name and stored as checkpoint metadata.
|
||||||
|
* ``CONDITIONING`` output: common text input fields are checked for
|
||||||
|
prompt text and stored as prompt metadata.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Input field names that carry a model path in loader-style nodes.
|
||||||
|
_MODEL_NAME_FIELDS = (
|
||||||
|
"ckpt_name", "unet_name", "model_path", "model_name", "gguf_name",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Extensions used by checkpoint_scanner.py — only record values that look
|
||||||
|
# like real model filenames to avoid capturing unrelated string fields.
|
||||||
|
_MODEL_EXTENSIONS = {
|
||||||
|
".ckpt", ".pt", ".pt2", ".bin", ".pth", ".safetensors", ".pkl", ".sft", ".gguf",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Input field names that may carry prompt text in encoder-style nodes.
|
||||||
|
_TEXT_FIELDS = ("text", "clip_l", "t5xxl", "prompt", "positive", "negative")
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def extract(node_id, inputs, outputs, metadata):
|
def extract(node_id, inputs, outputs, metadata, return_types=None):
|
||||||
pass
|
if return_types is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
# — MODEL loader detection (checkpoint / UNET / GGUF) —
|
||||||
|
if "MODEL" in return_types or any("MODEL" in str(t) for t in return_types):
|
||||||
|
for field in GenericNodeExtractor._MODEL_NAME_FIELDS:
|
||||||
|
val = inputs.get(field)
|
||||||
|
if val and isinstance(val, str) and val.strip():
|
||||||
|
name = val.strip()
|
||||||
|
if not any(name.lower().endswith(ext) for ext in GenericNodeExtractor._MODEL_EXTENSIONS):
|
||||||
|
continue
|
||||||
|
_store_checkpoint_metadata(metadata, node_id, name)
|
||||||
|
return
|
||||||
|
|
||||||
|
# — CONDITIONING encoder detection (CLIPTextEncode, Flux, custom) —
|
||||||
|
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
|
||||||
|
if text:
|
||||||
|
prompt_data = metadata.setdefault(PROMPTS, {})
|
||||||
|
prompt_data[node_id] = {
|
||||||
|
"text": text,
|
||||||
|
"node_id": node_id,
|
||||||
|
}
|
||||||
|
|
||||||
|
@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
|
||||||
|
if outputs and isinstance(outputs, list) and len(outputs) > 0:
|
||||||
|
if isinstance(outputs[0], tuple) and len(outputs[0]) > 0:
|
||||||
|
cond = outputs[0][0]
|
||||||
|
if cond is not None:
|
||||||
|
metadata[PROMPTS][node_id]["conditioning"] = cond
|
||||||
|
|
||||||
class CheckpointLoaderExtractor(NodeMetadataExtractor):
|
class CheckpointLoaderExtractor(NodeMetadataExtractor):
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -1154,6 +1221,35 @@ class CR_ApplyControlNetStackExtractor(NodeMetadataExtractor):
|
|||||||
metadata[PROMPTS][node_id]["positive_encoded"] = transformed_positive
|
metadata[PROMPTS][node_id]["positive_encoded"] = transformed_positive
|
||||||
metadata[PROMPTS][node_id]["negative_encoded"] = transformed_negative
|
metadata[PROMPTS][node_id]["negative_encoded"] = transformed_negative
|
||||||
|
|
||||||
|
class MetadataOverwriteExtractor(NodeMetadataExtractor):
|
||||||
|
"""Extract manually specified metadata from MetadataOverwriteLM node.
|
||||||
|
|
||||||
|
Stores truthy input values under the OVERWRITE category so that
|
||||||
|
extract_generation_params can merge them over the inferred params.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def extract(node_id, inputs, outputs, metadata):
|
||||||
|
if not inputs:
|
||||||
|
return
|
||||||
|
|
||||||
|
overwrite_params = {}
|
||||||
|
for key in METADATA_OVERWRITE_FIELDS:
|
||||||
|
value = inputs.get(key)
|
||||||
|
if key == "clip_skip":
|
||||||
|
if value != CLIP_SKIP_SENTINEL:
|
||||||
|
overwrite_params[key] = value
|
||||||
|
elif value: # truthy — only overwrite when user provided a real value
|
||||||
|
overwrite_params[key] = value
|
||||||
|
|
||||||
|
if overwrite_params:
|
||||||
|
metadata.setdefault(OVERWRITE, {})
|
||||||
|
metadata[OVERWRITE][node_id] = {
|
||||||
|
"parameters": overwrite_params,
|
||||||
|
"node_id": node_id,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
# Registry of node-specific extractors
|
# Registry of node-specific extractors
|
||||||
# Keys are node class names
|
# Keys are node class names
|
||||||
NODE_EXTRACTORS = {
|
NODE_EXTRACTORS = {
|
||||||
@@ -1221,5 +1317,7 @@ NODE_EXTRACTORS = {
|
|||||||
"CFGGuider": CFGGuiderExtractor, # Add CFGGuider
|
"CFGGuider": CFGGuiderExtractor, # Add CFGGuider
|
||||||
# Image
|
# Image
|
||||||
"VAEDecode": VAEDecodeExtractor, # Added VAEDecode extractor
|
"VAEDecode": VAEDecodeExtractor, # Added VAEDecode extractor
|
||||||
|
# Metadata overwrite
|
||||||
|
"MetadataOverwriteLM": MetadataOverwriteExtractor,
|
||||||
# Add other nodes as needed
|
# Add other nodes as needed
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,233 @@
|
|||||||
|
"""Metadata operations — thin in-process wrappers around LoRA Manager internal services.
|
||||||
|
|
||||||
|
All functions are simple Python async functions that delegate to the
|
||||||
|
appropriate internal service. They use **relative imports** within the
|
||||||
|
``py`` package, so ``sys.modules`` caching works normally and there is no
|
||||||
|
risk of double import or circular dependencies.
|
||||||
|
|
||||||
|
Usage (in-process, primary)::
|
||||||
|
|
||||||
|
from py.metadata_ops import list_base_models, read_metadata
|
||||||
|
|
||||||
|
models = await list_base_models()
|
||||||
|
meta = await read_metadata("/path/to/model.safetensors")
|
||||||
|
|
||||||
|
Usage (subprocess, debugging / external)::
|
||||||
|
|
||||||
|
python -m py.metadata_ops base-models list
|
||||||
|
python -m py.metadata_ops metadata read /path/to/model.safetensors
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Helpers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
SCANNER_TYPE_MAP: dict[str, str] = {
|
||||||
|
"get_lora_scanner": "lora",
|
||||||
|
"get_checkpoint_scanner": "checkpoint",
|
||||||
|
"get_embedding_scanner": "embedding",
|
||||||
|
}
|
||||||
|
|
||||||
|
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]:
|
||||||
|
"""Iterate all scanners and return the first (scanner, entry, getter_name)
|
||||||
|
that owns *model_path*. Returns ``(None, None, None)`` when no scanner
|
||||||
|
claims it.
|
||||||
|
"""
|
||||||
|
from ..services.service_registry import ServiceRegistry
|
||||||
|
|
||||||
|
normalized = os.path.normpath(model_path)
|
||||||
|
for getter_name in SCANNER_GETTER_NAMES:
|
||||||
|
getter = getattr(ServiceRegistry, getter_name, None)
|
||||||
|
if getter is None:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
scanner = await getter()
|
||||||
|
if scanner is None:
|
||||||
|
continue
|
||||||
|
cache = await scanner.get_cached_data()
|
||||||
|
for entry in cache.raw_data:
|
||||||
|
if os.path.normpath(entry.get("file_path", "")) == normalized:
|
||||||
|
return scanner, entry, getter_name
|
||||||
|
except Exception as exc:
|
||||||
|
logger.debug(
|
||||||
|
"Scanner %s check failed for %s: %s",
|
||||||
|
getter_name, model_path, exc,
|
||||||
|
)
|
||||||
|
return None, None, None
|
||||||
|
|
||||||
|
|
||||||
|
async def _find_scanner_for_model(
|
||||||
|
model_path: str,
|
||||||
|
) -> tuple[object, object] | tuple[None, None]:
|
||||||
|
"""Find the (scanner, cache_entry) responsible for *model_path*."""
|
||||||
|
scanner, entry, _ = await _find_model_entry(model_path)
|
||||||
|
return scanner, entry
|
||||||
|
|
||||||
|
|
||||||
|
async def identify_model_type(model_path: str) -> str:
|
||||||
|
"""Determine the model type (``\"lora\"``, ``\"checkpoint\"``, or
|
||||||
|
``\"embedding\"``) for *model_path*.
|
||||||
|
|
||||||
|
Falls back to ``\"lora\"`` when unknown.
|
||||||
|
"""
|
||||||
|
_, _, getter_name = await _find_model_entry(model_path)
|
||||||
|
return SCANNER_TYPE_MAP[getter_name] if getter_name else "lora"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Public API
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def list_base_models(limit: int = 0) -> List[str]:
|
||||||
|
"""Return all valid CivitAI base model names.
|
||||||
|
|
||||||
|
Uses ``CivitaiBaseModelService.get_base_models()`` which merges a
|
||||||
|
hardcoded list (``SUPPORTED_DOWNLOAD_SKIP_BASE_MODELS``) with remote
|
||||||
|
models fetched from the CivitAI API. Never empty — the hardcoded
|
||||||
|
fallback always provides a complete set.
|
||||||
|
|
||||||
|
The result is sorted alphabetically. Pass *limit* = 0 for all models.
|
||||||
|
"""
|
||||||
|
from ..services.civitai_base_model_service import (
|
||||||
|
CivitaiBaseModelService,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
service = await CivitaiBaseModelService.get_instance()
|
||||||
|
response = await service.get_base_models()
|
||||||
|
names: List[str] = response.get("models", [])
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("list_base_models failed: %s", exc)
|
||||||
|
names = []
|
||||||
|
if limit > 0:
|
||||||
|
return names[:limit]
|
||||||
|
return names
|
||||||
|
|
||||||
|
|
||||||
|
async def read_metadata(model_path: str) -> Dict[str, Any]:
|
||||||
|
"""Load the full metadata payload for *model_path* from disk.
|
||||||
|
|
||||||
|
Returns an empty dict when the metadata file does not exist or cannot
|
||||||
|
be parsed — never raises.
|
||||||
|
"""
|
||||||
|
from ..utils.metadata_manager import MetadataManager
|
||||||
|
|
||||||
|
try:
|
||||||
|
return await MetadataManager.load_metadata_payload(model_path) or {}
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("read_metadata failed for %s: %s", model_path, exc)
|
||||||
|
return {}
|
||||||
|
|
||||||
|
|
||||||
|
async def apply_metadata_updates(
|
||||||
|
model_path: str,
|
||||||
|
updates: Dict[str, Any],
|
||||||
|
) -> List[str]:
|
||||||
|
"""Merge *updates* into the model's on-disk metadata and persist.
|
||||||
|
|
||||||
|
Returns the list of field names that actually changed.
|
||||||
|
"""
|
||||||
|
from ..utils.metadata_manager import MetadataManager
|
||||||
|
|
||||||
|
metadata = await read_metadata(model_path)
|
||||||
|
updated_fields: List[str] = []
|
||||||
|
for key, value in updates.items():
|
||||||
|
old = metadata.get(key)
|
||||||
|
if old != value:
|
||||||
|
metadata[key] = value
|
||||||
|
updated_fields.append(key)
|
||||||
|
if updated_fields:
|
||||||
|
await MetadataManager.save_metadata(model_path, metadata)
|
||||||
|
return updated_fields
|
||||||
|
|
||||||
|
|
||||||
|
async def download_preview(
|
||||||
|
model_path: str,
|
||||||
|
url: str,
|
||||||
|
*,
|
||||||
|
target_width: int = 480,
|
||||||
|
quality: int = 85,
|
||||||
|
) -> str | None:
|
||||||
|
"""Download a preview image from *url*, optimise to .webp, and save it.
|
||||||
|
|
||||||
|
The output file is placed alongside the model file with a ``.webp``
|
||||||
|
extension. Returns the local file path on success, ``None`` on failure.
|
||||||
|
"""
|
||||||
|
from ..services.downloader import get_downloader
|
||||||
|
from ..utils.exif_utils import ExifUtils
|
||||||
|
|
||||||
|
if not url or not url.strip():
|
||||||
|
return None
|
||||||
|
|
||||||
|
base_name = os.path.splitext(os.path.basename(model_path))[0]
|
||||||
|
preview_dir = os.path.dirname(model_path)
|
||||||
|
output_path = os.path.join(preview_dir, base_name + ".webp")
|
||||||
|
|
||||||
|
downloader = await get_downloader()
|
||||||
|
|
||||||
|
# Try in-memory download + optimise first
|
||||||
|
success, content, _headers = await downloader.download_to_memory(
|
||||||
|
url, use_auth=False,
|
||||||
|
)
|
||||||
|
if success and content:
|
||||||
|
try:
|
||||||
|
optimized_data, _ = ExifUtils.optimize_image(
|
||||||
|
image_data=content,
|
||||||
|
target_width=target_width,
|
||||||
|
format="webp",
|
||||||
|
quality=quality,
|
||||||
|
preserve_metadata=False,
|
||||||
|
)
|
||||||
|
with open(output_path, "wb") as f:
|
||||||
|
f.write(optimized_data)
|
||||||
|
return output_path
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Preview optimisation failed, saving raw: %s", exc)
|
||||||
|
# Fall through to raw save
|
||||||
|
|
||||||
|
# Fallback: download directly to file
|
||||||
|
try:
|
||||||
|
ok, _ = await downloader.download_file(url, output_path, use_auth=False)
|
||||||
|
if ok:
|
||||||
|
return output_path
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Preview fallback download failed for %s: %s", model_path, exc)
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
async def refresh_cache(model_path: str) -> bool:
|
||||||
|
"""Invalidate and reload the scanner cache entry for *model_path*.
|
||||||
|
|
||||||
|
Returns ``True`` when the model was found and the cache was refreshed.
|
||||||
|
"""
|
||||||
|
scanner, entry = await _find_scanner_for_model(model_path)
|
||||||
|
if scanner is None:
|
||||||
|
logger.warning("refresh_cache: no scanner found for %s", model_path)
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
metadata = await read_metadata(model_path)
|
||||||
|
if not metadata:
|
||||||
|
logger.warning("refresh_cache: no metadata for %s", model_path)
|
||||||
|
return False
|
||||||
|
await scanner.update_single_model_cache(model_path, model_path, metadata)
|
||||||
|
return True
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("refresh_cache failed for %s: %s", model_path, exc)
|
||||||
|
return False
|
||||||
@@ -0,0 +1,113 @@
|
|||||||
|
"""Subprocess entry point for ``metadata_ops`` (debugging / external use).
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
python -m py.metadata_ops base-models list [--limit N]
|
||||||
|
python -m py.metadata_ops metadata read <path>
|
||||||
|
python -m py.metadata_ops metadata update <path> --json '{...}'
|
||||||
|
python -m py.metadata_ops preview download <path> --url <url>
|
||||||
|
python -m py.metadata_ops cache refresh <path>
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import sys
|
||||||
|
from typing import Any, Dict, List
|
||||||
|
|
||||||
|
|
||||||
|
def _build_parser() -> argparse.ArgumentParser:
|
||||||
|
parser = argparse.ArgumentParser(prog="lmcli", description="LoRA Manager Agent CLI")
|
||||||
|
sub = parser.add_subparsers(dest="command", required=True)
|
||||||
|
|
||||||
|
# base-models list
|
||||||
|
base_models = sub.add_parser("base-models", aliases=["bm"])
|
||||||
|
base_models_cmds = base_models.add_subparsers(dest="subcommand", required=True)
|
||||||
|
base_models_list = base_models_cmds.add_parser("list")
|
||||||
|
base_models_list.add_argument(
|
||||||
|
"--limit", type=int, default=0, help="Max number of models (0 = all)"
|
||||||
|
)
|
||||||
|
|
||||||
|
# metadata read
|
||||||
|
meta = sub.add_parser("metadata", aliases=["md"])
|
||||||
|
meta_cmds = meta.add_subparsers(dest="subcommand", required=True)
|
||||||
|
meta_read = meta_cmds.add_parser("read")
|
||||||
|
meta_read.add_argument("path", type=str, help="Model file path")
|
||||||
|
|
||||||
|
# metadata update
|
||||||
|
meta_update = meta_cmds.add_parser("update")
|
||||||
|
meta_update.add_argument("path", type=str, help="Model file path")
|
||||||
|
meta_update.add_argument(
|
||||||
|
"--json",
|
||||||
|
type=str,
|
||||||
|
required=True,
|
||||||
|
help='JSON object of fields to update, e.g. \'{"base_model": "SDXL 1.0"}\'',
|
||||||
|
)
|
||||||
|
|
||||||
|
# preview download
|
||||||
|
prev = sub.add_parser("preview", aliases=["pv"])
|
||||||
|
prev_cmds = prev.add_subparsers(dest="subcommand", required=True)
|
||||||
|
prev_dl = prev_cmds.add_parser("download")
|
||||||
|
prev_dl.add_argument("path", type=str, help="Model file path")
|
||||||
|
prev_dl.add_argument("--url", type=str, required=True, help="Preview image URL")
|
||||||
|
|
||||||
|
# cache refresh
|
||||||
|
cache = sub.add_parser("cache")
|
||||||
|
cache_cmds = cache.add_subparsers(dest="subcommand", required=True)
|
||||||
|
cache_refresh = cache_cmds.add_parser("refresh")
|
||||||
|
cache_refresh.add_argument("path", type=str, help="Model file path")
|
||||||
|
|
||||||
|
return parser
|
||||||
|
|
||||||
|
|
||||||
|
async def _run(args: argparse.Namespace) -> Any:
|
||||||
|
from . import ( # lazy import so startup is fast
|
||||||
|
list_base_models,
|
||||||
|
read_metadata,
|
||||||
|
apply_metadata_updates,
|
||||||
|
download_preview,
|
||||||
|
refresh_cache,
|
||||||
|
)
|
||||||
|
|
||||||
|
cmd = args.command
|
||||||
|
sub = args.subcommand
|
||||||
|
|
||||||
|
if cmd in ("base-models", "bm") and sub == "list":
|
||||||
|
return await list_base_models(limit=args.limit)
|
||||||
|
|
||||||
|
if cmd in ("metadata", "md") and sub == "read":
|
||||||
|
return await read_metadata(args.path)
|
||||||
|
|
||||||
|
if cmd in ("metadata", "md") and sub == "update":
|
||||||
|
updates: Dict[str, Any] = json.loads(args.json)
|
||||||
|
return await apply_metadata_updates(args.path, updates)
|
||||||
|
|
||||||
|
if cmd in ("preview", "pv") and sub == "download":
|
||||||
|
return await download_preview(args.path, args.url)
|
||||||
|
|
||||||
|
if cmd == "cache" and sub == "refresh":
|
||||||
|
return await refresh_cache(args.path)
|
||||||
|
|
||||||
|
raise ValueError(f"Unknown command: {cmd} {sub}")
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
parser = _build_parser()
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
result = asyncio.run(_run(args))
|
||||||
|
# Always print as JSON so callers can parse reliably
|
||||||
|
if isinstance(result, list):
|
||||||
|
for item in result:
|
||||||
|
print(item)
|
||||||
|
elif isinstance(result, dict):
|
||||||
|
json.dump(result, sys.stdout, ensure_ascii=False, indent=2)
|
||||||
|
print()
|
||||||
|
else:
|
||||||
|
print(json.dumps(result))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -41,7 +41,12 @@ async def api_json_error(
|
|||||||
if exc.status < 400:
|
if exc.status < 400:
|
||||||
raise
|
raise
|
||||||
|
|
||||||
logger.warning(
|
# Preview 404 is routine (file deleted from disk) — not worth a warning.
|
||||||
|
logger_method = logger.warning
|
||||||
|
if request.path.startswith("/api/lm/previews") and exc.status == 404:
|
||||||
|
logger_method = logger.debug
|
||||||
|
|
||||||
|
logger_method(
|
||||||
"API %s %s returned HTTP %d: %s",
|
"API %s %s returned HTTP %d: %s",
|
||||||
request.method,
|
request.method,
|
||||||
request.path,
|
request.path,
|
||||||
|
|||||||
@@ -0,0 +1,117 @@
|
|||||||
|
"""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,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class CreateHookLoraLM:
|
||||||
|
NAME = "Create Hook LoRA (LoraManager)"
|
||||||
|
CATEGORY = "Lora Manager/hooks"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"text": (
|
||||||
|
"AUTOCOMPLETE_TEXT_LORAS",
|
||||||
|
{
|
||||||
|
"placeholder": "Search LoRAs to add...",
|
||||||
|
"tooltip": (
|
||||||
|
"Search and select LoRAs. Each LoRA gets its own "
|
||||||
|
"model/clip strength. Hooks chain with prev_hooks."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
),
|
||||||
|
},
|
||||||
|
"optional": FlexibleOptionalInputType(any_type),
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("HOOKS", "STRING", "STRING")
|
||||||
|
RETURN_NAMES = ("HOOKS", "trigger_words", "active_loras")
|
||||||
|
FUNCTION = "create_hook"
|
||||||
|
|
||||||
|
def create_hook(self, text: str, **kwargs):
|
||||||
|
"""Create a HookGroup from the selected LoRAs, chained with prev_hooks.
|
||||||
|
|
||||||
|
Each active LoRA from the widget is loaded and wrapped in a WeightHook
|
||||||
|
via :func:`comfy.hooks.create_hook_lora`. All hooks are combined into a
|
||||||
|
single group and returned alongside trigger words and a human-readable
|
||||||
|
summary of the active LoRAs.
|
||||||
|
"""
|
||||||
|
del text # used by the frontend widget only
|
||||||
|
|
||||||
|
# Lazy imports: comfy is not available in CI/test environment at module level
|
||||||
|
import comfy.hooks # type: ignore # noqa: C0415
|
||||||
|
import comfy.utils # type: ignore # noqa: C0415
|
||||||
|
|
||||||
|
prev_hooks: comfy.hooks.HookGroup | None = kwargs.get("prev_hooks")
|
||||||
|
|
||||||
|
hook_group = prev_hooks.clone() if prev_hooks is not None else comfy.hooks.HookGroup()
|
||||||
|
|
||||||
|
all_trigger_words: list[str] = []
|
||||||
|
active_loras: list[tuple[str, float, float]] = []
|
||||||
|
|
||||||
|
for lora in get_loras_list(kwargs):
|
||||||
|
if not lora.get("active", False):
|
||||||
|
continue
|
||||||
|
|
||||||
|
lora_name = apply_lora_syntax_format(lora["name"])
|
||||||
|
model_strength = float(lora["strength"])
|
||||||
|
clip_strength = float(lora.get("clipStrength", model_strength))
|
||||||
|
|
||||||
|
# Skip useless no-op entries (both strengths are zero)
|
||||||
|
if model_strength == 0.0 and clip_strength == 0.0:
|
||||||
|
continue
|
||||||
|
|
||||||
|
lora_path, trigger_words = get_lora_info_absolute(lora_name)
|
||||||
|
if not lora_path or not os.path.isfile(lora_path):
|
||||||
|
logger.warning("LoRA '%s' not found — skipping", lora_name)
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
lora_weights = comfy.utils.load_torch_file(lora_path, safe_load=True)
|
||||||
|
|
||||||
|
lora_hooks = comfy.hooks.create_hook_lora(
|
||||||
|
lora=lora_weights,
|
||||||
|
strength_model=model_strength,
|
||||||
|
strength_clip=clip_strength,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Failed to load LoRA '%s' — skipping", lora_name)
|
||||||
|
continue
|
||||||
|
hook_group = hook_group.clone_and_combine(lora_hooks)
|
||||||
|
|
||||||
|
active_loras.append((lora_name, model_strength, clip_strength))
|
||||||
|
all_trigger_words.extend(trigger_words)
|
||||||
|
|
||||||
|
# Format trigger words (group mode separator)
|
||||||
|
trigger_words_text = ",, ".join(all_trigger_words) if all_trigger_words else ""
|
||||||
|
|
||||||
|
# Format active LoRAs summary
|
||||||
|
formatted_loras = []
|
||||||
|
for name, model_s, clip_s in active_loras:
|
||||||
|
if abs(model_s - clip_s) > 0.001:
|
||||||
|
formatted_loras.append(
|
||||||
|
f"<lora:{name}:{model_s}:{clip_s}>"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
formatted_loras.append(f"<lora:{name}:{model_s}>")
|
||||||
|
active_loras_text = " ".join(formatted_loras)
|
||||||
|
|
||||||
|
return (hook_group, trigger_words_text, active_loras_text)
|
||||||
@@ -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)",
|
||||||
|
}
|
||||||
+2
-17
@@ -1,6 +1,5 @@
|
|||||||
import importlib
|
import importlib
|
||||||
import logging
|
import logging
|
||||||
import re
|
|
||||||
|
|
||||||
import comfy.sd # type: ignore
|
import comfy.sd # type: ignore
|
||||||
import comfy.utils # type: ignore
|
import comfy.utils # type: ignore
|
||||||
@@ -14,6 +13,7 @@ from .utils import (
|
|||||||
extract_lora_name,
|
extract_lora_name,
|
||||||
get_loras_list,
|
get_loras_list,
|
||||||
nunchaku_load_lora,
|
nunchaku_load_lora,
|
||||||
|
parse_lora_syntax,
|
||||||
)
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -189,25 +189,10 @@ class LoraTextLoaderLM:
|
|||||||
RETURN_NAMES = ("MODEL", "CLIP", "trigger_words", "loaded_loras")
|
RETURN_NAMES = ("MODEL", "CLIP", "trigger_words", "loaded_loras")
|
||||||
FUNCTION = "load_loras_from_text"
|
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):
|
def load_loras_from_text(self, model, lora_syntax, clip=None, lora_stack=None):
|
||||||
"""Load LoRAs based on text syntax input."""
|
"""Load LoRAs based on text syntax input."""
|
||||||
lora_entries = _collect_stack_entries(lora_stack)
|
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_path, trigger_words = get_lora_info_absolute(lora["name"])
|
||||||
lora_entries.append({
|
lora_entries.append({
|
||||||
"name": lora["name"],
|
"name": lora["name"],
|
||||||
|
|||||||
@@ -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,170 @@
|
|||||||
|
"""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,
|
||||||
|
METADATA_OVERWRITE_FIELDS,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
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",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"tooltip": "Sampler name. Only overwrites when non-empty.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"scheduler": (
|
||||||
|
"STRING",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"tooltip": "Scheduler name. Only overwrites when non-empty.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"model": (
|
||||||
|
"STRING",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"tooltip": (
|
||||||
|
"The checkpoint or diffusion model (UNet) used "
|
||||||
|
"for generation. 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.
|
||||||
|
"""
|
||||||
|
result: dict[str, Any] = {}
|
||||||
|
for key in METADATA_OVERWRITE_FIELDS:
|
||||||
|
value = kwargs.get(key)
|
||||||
|
if key == "clip_skip":
|
||||||
|
if value != _CLIP_SKIP_SENTINEL:
|
||||||
|
result[key] = value
|
||||||
|
elif value:
|
||||||
|
result[key] = value
|
||||||
|
return (result,)
|
||||||
+346
-127
@@ -16,6 +16,156 @@ from PIL import Image, PngImagePlugin
|
|||||||
import piexif
|
import piexif
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
|
# Civitai-compatible sampler name mapping: ComfyUI internal → A1111 display name
|
||||||
|
CIVITAI_SAMPLER_MAP = {
|
||||||
|
"euler": "Euler",
|
||||||
|
"euler_ancestral": "Euler a",
|
||||||
|
"lms": "LMS",
|
||||||
|
"heun": "Heun",
|
||||||
|
"dpm_2": "DPM2",
|
||||||
|
"dpm_2_ancestral": "DPM2 a",
|
||||||
|
"dpmpp_2s_ancestral": "DPM++ 2S a",
|
||||||
|
"dpmpp_2m": "DPM++ 2M",
|
||||||
|
"dpmpp_sde": "DPM++ SDE",
|
||||||
|
"dpmpp_sde_gpu": "DPM++ SDE",
|
||||||
|
"dpmpp_2m_sde": "DPM++ 2M SDE",
|
||||||
|
"dpmpp_2m_sde_gpu": "DPM++ 2M SDE",
|
||||||
|
"dpmpp_3m_sde": "DPM++ 3M SDE",
|
||||||
|
"dpm_fast": "DPM fast",
|
||||||
|
"dpm_adaptive": "DPM adaptive",
|
||||||
|
"ddim": "DDIM",
|
||||||
|
"plms": "PLMS",
|
||||||
|
"uni_pc_bh2": "UniPC",
|
||||||
|
"uni_pc": "UniPC",
|
||||||
|
"lcm": "LCM",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Base model display name → AIR URN slug
|
||||||
|
# Sourced from civitai source: src/shared/constants/basemodel.constants.ts
|
||||||
|
BASE_MODEL_AIR_SLUG = {
|
||||||
|
# Stable Diffusion family
|
||||||
|
"SD 1.4": "sd1",
|
||||||
|
"SD 1.5": "sd1",
|
||||||
|
"SD 1.5 LCM": "sd1",
|
||||||
|
"SD 1.5 Hyper": "sd1",
|
||||||
|
"SD 2.0": "sd2",
|
||||||
|
"SD 2.0 768": "sd2",
|
||||||
|
"SD 2.1": "sd2",
|
||||||
|
"SD 2.1 768": "sd2",
|
||||||
|
"SD 2.1 Unclip": "sd2",
|
||||||
|
"SD 3.0": "sd3",
|
||||||
|
"SD 3.5": "sd35",
|
||||||
|
"SD 3.5 Large": "sd35",
|
||||||
|
"SD 3.5 Large Turbo": "sd35",
|
||||||
|
"SD 3.5 Medium": "sd35",
|
||||||
|
"SDXL 0.9": "sdxl",
|
||||||
|
"SDXL 1.0": "sdxl",
|
||||||
|
"SDXL 1.0 LCM": "sdxl",
|
||||||
|
"SDXL Lightning": "sdxl",
|
||||||
|
"SDXL Hyper": "sdxl",
|
||||||
|
"SDXL Turbo": "sdxl",
|
||||||
|
"SDXL Distilled": "sdxldistilled",
|
||||||
|
"Stable Cascade": "scascade",
|
||||||
|
"Stable Video Diffusion": "svd",
|
||||||
|
"SVD": "svd",
|
||||||
|
"SVD XT": "svdxt",
|
||||||
|
|
||||||
|
# SDXL community fine-tunes
|
||||||
|
"Pony": "pony",
|
||||||
|
"Pony Diffusion": "pony",
|
||||||
|
"Illustrious": "illustrious",
|
||||||
|
"NoobAI": "noobai",
|
||||||
|
"Animagine": "illustrious",
|
||||||
|
|
||||||
|
# Flux family
|
||||||
|
"Flux.1": "flux1",
|
||||||
|
"Flux.1 D": "flux1",
|
||||||
|
"Flux.1 S": "flux1",
|
||||||
|
"Flux.1 Krea": "fluxkrea",
|
||||||
|
"Flux.1 Kontext": "flux1kontext",
|
||||||
|
"Flux.2": "flux2",
|
||||||
|
"Flux.2 D": "flux2",
|
||||||
|
"Flux.2 Klein 9B": "flux2klein_9b",
|
||||||
|
"Flux.2 Klein 9B Base": "flux2klein_9b_base",
|
||||||
|
"Flux.2 Klein 4B": "flux2klein_4b",
|
||||||
|
"Flux.2 Klein 4B Base": "flux2klein_4b_base",
|
||||||
|
|
||||||
|
# Other image models (sorted alphabetically)
|
||||||
|
"AuraFlow": "auraflow",
|
||||||
|
"Chroma": "chroma",
|
||||||
|
"HiDream": "hidream",
|
||||||
|
"HiDream-O1": "hidream-o1",
|
||||||
|
"Hunyuan DiT": "hydit1",
|
||||||
|
"Hunyuan Video": "hyv1",
|
||||||
|
"Kolors": "kolors",
|
||||||
|
"Lumina": "lumina",
|
||||||
|
"Mochi": "mochi",
|
||||||
|
"ODOR": "odor",
|
||||||
|
"PixArt Alpha": "pixarta",
|
||||||
|
"PixArt Sigma": "pixarte",
|
||||||
|
"Playground v2": "playgroundv2",
|
||||||
|
"Playground v2.5": "playgroundv2",
|
||||||
|
"Pony Diffusion V7": "ponyv7",
|
||||||
|
|
||||||
|
# Video models
|
||||||
|
"CogVideoX": "cogvideox",
|
||||||
|
"LTX Video": "ltxv",
|
||||||
|
"LTX Video 2": "ltxv2",
|
||||||
|
"LTX Video 2.3": "ltxv23",
|
||||||
|
"Wan Video": "wanvideo",
|
||||||
|
"Wan Video 1.3B T2V": "wanvideo_13b_t2v",
|
||||||
|
"Wan Video 14B T2V": "wanvideo_14b_t2v",
|
||||||
|
"Wan Video 14B I2V 480p": "wanvideo_14b_i2v_480p",
|
||||||
|
"Wan Video 14B I2V 720p": "wanvideo_14b_i2v_720p",
|
||||||
|
|
||||||
|
# Third-party / proprietary image models
|
||||||
|
"Boogu": "boogu",
|
||||||
|
"Ernie": "ernie",
|
||||||
|
"Grok": "grok",
|
||||||
|
"HappyHorse": "happyhorse",
|
||||||
|
"Ideogram": "ideogram",
|
||||||
|
"Ideogram 4.0": "ideogram",
|
||||||
|
"Imagen": "imagen4",
|
||||||
|
"Imagen 4": "imagen4",
|
||||||
|
"Krea": "krea2",
|
||||||
|
"Krea 2": "krea2",
|
||||||
|
"Lens": "lens",
|
||||||
|
"MAI": "mai",
|
||||||
|
"Nano Banana": "nanobanana",
|
||||||
|
"OpenAI": "openai",
|
||||||
|
"Reve": "reve",
|
||||||
|
"Reve 2": "reve",
|
||||||
|
"Reve 2.1": "reve",
|
||||||
|
"Seedream": "seedream",
|
||||||
|
"Sora": "sora2",
|
||||||
|
"Sora 2": "sora2",
|
||||||
|
"Veo": "veo3",
|
||||||
|
"Veo 2": "veo3",
|
||||||
|
"Veo 3": "veo3",
|
||||||
|
"ZImageTurbo": "zimageturbo",
|
||||||
|
"ZImageBase": "zimagebase",
|
||||||
|
"ZImage": "zimagebase",
|
||||||
|
|
||||||
|
# Third-party video models
|
||||||
|
"Hailuo by MiniMax": "minimax",
|
||||||
|
"Haiper": "haiper",
|
||||||
|
"Kling": "kling",
|
||||||
|
"Lightricks": "lightricks",
|
||||||
|
"Seedance": "seedance",
|
||||||
|
"Vidu": "vidu",
|
||||||
|
|
||||||
|
# Qwen family
|
||||||
|
"Qwen": "qwen",
|
||||||
|
"Qwen 2": "qwen2",
|
||||||
|
|
||||||
|
# Anima
|
||||||
|
"Anima": "anima",
|
||||||
|
|
||||||
|
# Special
|
||||||
|
"Upscaler": "upscaler",
|
||||||
|
"Other": "other",
|
||||||
|
}
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@@ -70,11 +220,29 @@ class SaveImageLM:
|
|||||||
"tooltip": "Compression quality for JPEG and lossy WebP formats (1-100). Higher values mean better quality but larger files.",
|
"tooltip": "Compression quality for JPEG and lossy WebP formats (1-100). Higher values mean better quality but larger files.",
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
|
"webp_method": (
|
||||||
|
"INT",
|
||||||
|
{
|
||||||
|
"default": 6,
|
||||||
|
"min": 0,
|
||||||
|
"max": 6,
|
||||||
|
"tooltip": "WebP compression method (0-6). 0=fastest/largest, 6=slowest/smallest. Only applies when file_format is 'webp'.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"jpeg_subsampling": (
|
||||||
|
"INT",
|
||||||
|
{
|
||||||
|
"default": 0,
|
||||||
|
"min": 0,
|
||||||
|
"max": 2,
|
||||||
|
"tooltip": "JPEG chroma subsampling level. 0=4:4:4 (best quality), 1=4:2:2, 2=4:2:0 (smallest files). Only applies when file_format is 'jpeg'.",
|
||||||
|
},
|
||||||
|
),
|
||||||
"embed_workflow": (
|
"embed_workflow": (
|
||||||
"BOOLEAN",
|
"BOOLEAN",
|
||||||
{
|
{
|
||||||
"default": False,
|
"default": False,
|
||||||
"tooltip": "Embeds the complete workflow data into the image metadata. Only works with PNG and WebP formats.",
|
"tooltip": "When enabled, saved images store the complete workflow. Drag the image back into ComfyUI to restore the original node graph. PNG and WebP only.",
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
"save_with_metadata": (
|
"save_with_metadata": (
|
||||||
@@ -142,148 +310,194 @@ class SaveImageLM:
|
|||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def format_metadata(self, metadata_dict):
|
def _resolve_model_cache_entry(self, scanner_type: str, name: str):
|
||||||
"""Format metadata in the requested format similar to userComment example"""
|
"""Resolve model hash, civitai metadata, and base_model from scanner cache.
|
||||||
if not metadata_dict:
|
Returns (hash_str, civitai_dict, base_model_str). All values are empty defaults when not found."""
|
||||||
return ""
|
scanner = ServiceRegistry.get_service_sync(scanner_type)
|
||||||
|
if scanner is None or not name:
|
||||||
|
return "", {}, ""
|
||||||
|
|
||||||
# Helper function to only add parameter if value is not None
|
entry = self._get_cached_model_by_name(scanner, name)
|
||||||
def add_param_if_not_none(param_list, label, value):
|
if entry is None:
|
||||||
if value is not None:
|
basename = os.path.splitext(os.path.basename(name))[0]
|
||||||
param_list.append(f"{label}: {value}")
|
hash_val = scanner.get_hash_by_filename(basename)
|
||||||
|
return (hash_val or "").lower(), {}, ""
|
||||||
|
|
||||||
|
hash_val = (entry.get("sha256") or "").lower()
|
||||||
|
civitai = entry.get("civitai") or {}
|
||||||
|
base_model = entry.get("base_model") or ""
|
||||||
|
return hash_val, civitai, base_model
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_civitai_sampler_name(sampler_name: str, scheduler: str) -> str:
|
||||||
|
if sampler_name in CIVITAI_SAMPLER_MAP:
|
||||||
|
civitai_name = CIVITAI_SAMPLER_MAP[sampler_name]
|
||||||
|
if scheduler == "karras":
|
||||||
|
civitai_name += " Karras"
|
||||||
|
elif scheduler == "exponential":
|
||||||
|
civitai_name += " Exponential"
|
||||||
|
return civitai_name
|
||||||
|
else:
|
||||||
|
if scheduler and scheduler != "normal":
|
||||||
|
return f"{sampler_name}_{scheduler}"
|
||||||
|
return sampler_name
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _build_air_string(base_model: str, model_type: str, model_id: int, version_id: int) -> str:
|
||||||
|
slug = BASE_MODEL_AIR_SLUG.get(base_model, "other")
|
||||||
|
type_lower = model_type.lower() if model_type else "other"
|
||||||
|
return f"urn:air:{slug}:{type_lower}:civitai:{model_id}@{version_id}"
|
||||||
|
|
||||||
|
def format_metadata(self, metadata_dict: dict) -> str:
|
||||||
|
"""Format metadata as A1111-compatible parameters string with Hashes JSON and Civitai resources."""
|
||||||
|
if not metadata_dict: return ""
|
||||||
|
|
||||||
# Extract the prompt and negative prompt
|
|
||||||
prompt = metadata_dict.get("prompt", "")
|
prompt = metadata_dict.get("prompt", "")
|
||||||
negative_prompt = metadata_dict.get("negative_prompt", "")
|
negative_prompt = metadata_dict.get("negative_prompt", "")
|
||||||
|
steps = metadata_dict.get("steps")
|
||||||
# Extract loras from the prompt if present
|
cfg = metadata_dict.get("guidance")
|
||||||
|
if cfg is None:
|
||||||
|
cfg = metadata_dict.get("cfg_scale")
|
||||||
|
if cfg is None:
|
||||||
|
cfg = metadata_dict.get("cfg")
|
||||||
|
seed = metadata_dict.get("seed")
|
||||||
|
size = metadata_dict.get("size")
|
||||||
|
sampler = metadata_dict.get("sampler") or ""
|
||||||
|
scheduler = metadata_dict.get("scheduler") or "normal"
|
||||||
|
checkpoint = metadata_dict.get("checkpoint") or ""
|
||||||
loras_text = metadata_dict.get("loras", "")
|
loras_text = metadata_dict.get("loras", "")
|
||||||
lora_hashes = {}
|
clip_skip = metadata_dict.get("clip_skip")
|
||||||
|
|
||||||
# If loras are found, add them on a new line after the prompt
|
# Parse LoRA entries from <lora:name:strength> format
|
||||||
|
lora_entries: list[tuple[str, float]] = []
|
||||||
if loras_text:
|
if loras_text:
|
||||||
prompt_with_loras = f"{prompt}\n{loras_text}"
|
for match in re.findall(r"<lora:([^:]+):([^>]+)>", loras_text):
|
||||||
|
lora_name, strength_str = match
|
||||||
|
try:
|
||||||
|
strength = float(strength_str)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
strength = 1.0
|
||||||
|
lora_entries.append((lora_name, strength))
|
||||||
|
|
||||||
# Extract lora names from the format <lora:name:strength>
|
# Resolve checkpoint hash and Civitai data from local cache
|
||||||
lora_matches = re.findall(r"<lora:([^:]+):([^>]+)>", loras_text)
|
ckpt_hash, ckpt_civitai, ckpt_base_model = "", {}, ""
|
||||||
|
ckpt_display_name = ""
|
||||||
|
if checkpoint:
|
||||||
|
ckpt_hash, ckpt_civitai, ckpt_base_model = self._resolve_model_cache_entry(
|
||||||
|
"checkpoint_scanner", checkpoint
|
||||||
|
)
|
||||||
|
ckpt_display_name = os.path.splitext(os.path.basename(checkpoint))[0]
|
||||||
|
|
||||||
# Get hash for each lora
|
# Resolve LoRA hash and Civitai data from local cache
|
||||||
for lora_name, strength in lora_matches:
|
loras_data: list[dict] = []
|
||||||
hash_value = self.get_lora_hash(lora_name)
|
for lora_name, strength in lora_entries:
|
||||||
if hash_value:
|
lora_hash, lora_civitai, lora_base_model = self._resolve_model_cache_entry(
|
||||||
lora_hashes[lora_name] = hash_value
|
"lora_scanner", lora_name
|
||||||
else:
|
)
|
||||||
prompt_with_loras = prompt
|
loras_data.append({
|
||||||
|
"name": lora_name,
|
||||||
|
"strength": strength,
|
||||||
|
"hash": lora_hash,
|
||||||
|
"civitai": lora_civitai,
|
||||||
|
"base_model": lora_base_model,
|
||||||
|
})
|
||||||
|
|
||||||
# Format the first part (prompt and loras)
|
# Build Hashes JSON (A1111 / Civitai standard format)
|
||||||
metadata_parts = [prompt_with_loras]
|
hashes: dict[str, str] = {}
|
||||||
|
if ckpt_hash:
|
||||||
|
hashes["model"] = ckpt_hash[:10].upper()
|
||||||
|
for lora in loras_data:
|
||||||
|
if lora["hash"]:
|
||||||
|
hashes[f"LORA:{lora['name']}"] = lora["hash"][:10].upper()
|
||||||
|
|
||||||
# Add negative prompt
|
# Build Civitai resources JSON array
|
||||||
|
civitai_resources: list[dict] = []
|
||||||
|
if ckpt_civitai.get("id", 0) > 0:
|
||||||
|
ckpt_resource: dict = {}
|
||||||
|
ckpt_type = (ckpt_civitai.get("model") or {}).get("type", "Checkpoint")
|
||||||
|
model_id = ckpt_civitai.get("modelId", 0)
|
||||||
|
version_id = ckpt_civitai.get("id", 0)
|
||||||
|
if model_id and version_id:
|
||||||
|
ckpt_resource["air"] = self._build_air_string(
|
||||||
|
ckpt_base_model, ckpt_type, int(model_id), int(version_id)
|
||||||
|
)
|
||||||
|
elif version_id:
|
||||||
|
ckpt_resource["modelVersionId"] = int(version_id)
|
||||||
|
if ckpt_civitai.get("name"):
|
||||||
|
ckpt_resource["versionName"] = ckpt_civitai["name"]
|
||||||
|
if ckpt_resource:
|
||||||
|
civitai_resources.append(ckpt_resource)
|
||||||
|
|
||||||
|
for lora in loras_data:
|
||||||
|
lora_civitai = lora["civitai"]
|
||||||
|
if not lora_civitai or lora_civitai.get("id", 0) <= 0:
|
||||||
|
continue
|
||||||
|
lora_resource: dict = {"weight": lora["strength"]}
|
||||||
|
lora_type = (lora_civitai.get("model") or {}).get("type", "LORA")
|
||||||
|
model_id = lora_civitai.get("modelId", 0)
|
||||||
|
version_id = lora_civitai.get("id", 0)
|
||||||
|
if model_id and version_id:
|
||||||
|
lora_resource["air"] = self._build_air_string(
|
||||||
|
lora["base_model"], lora_type, int(model_id), int(version_id)
|
||||||
|
)
|
||||||
|
elif version_id:
|
||||||
|
lora_resource["modelVersionId"] = int(version_id)
|
||||||
|
if lora_civitai.get("name"):
|
||||||
|
lora_resource["versionName"] = lora_civitai["name"]
|
||||||
|
civitai_resources.append(lora_resource)
|
||||||
|
|
||||||
|
sampler_name = CIVITAI_SAMPLER_MAP.get(sampler, sampler) if sampler else None
|
||||||
|
|
||||||
|
scheduler_mapping = {
|
||||||
|
"normal": "Normal",
|
||||||
|
"karras": "Karras",
|
||||||
|
"exponential": "Exponential",
|
||||||
|
"sgm_uniform": "SGM Uniform",
|
||||||
|
"sgm_quadratic": "SGM Quadratic",
|
||||||
|
}
|
||||||
|
scheduler_name = scheduler_mapping.get(scheduler, scheduler) if scheduler else None
|
||||||
|
|
||||||
|
# Build output lines
|
||||||
|
lines = [prompt] if prompt else [""]
|
||||||
if negative_prompt:
|
if negative_prompt:
|
||||||
metadata_parts.append(f"Negative prompt: {negative_prompt}")
|
lines.append(f"Negative prompt: {negative_prompt}")
|
||||||
|
|
||||||
# Format the second part (generation parameters)
|
params: list[str] = []
|
||||||
params = []
|
if steps is not None:
|
||||||
|
params.append(f"Steps: {steps}")
|
||||||
# Add standard parameters in the correct order
|
|
||||||
if "steps" in metadata_dict:
|
|
||||||
add_param_if_not_none(params, "Steps", metadata_dict.get("steps"))
|
|
||||||
|
|
||||||
# Combine sampler and scheduler information
|
|
||||||
sampler_name = None
|
|
||||||
scheduler_name = None
|
|
||||||
|
|
||||||
if "sampler" in metadata_dict:
|
|
||||||
sampler = metadata_dict.get("sampler")
|
|
||||||
# Convert ComfyUI sampler names to user-friendly names
|
|
||||||
sampler_mapping = {
|
|
||||||
"euler": "Euler",
|
|
||||||
"euler_ancestral": "Euler a",
|
|
||||||
"dpm_2": "DPM2",
|
|
||||||
"dpm_2_ancestral": "DPM2 a",
|
|
||||||
"heun": "Heun",
|
|
||||||
"dpm_fast": "DPM fast",
|
|
||||||
"dpm_adaptive": "DPM adaptive",
|
|
||||||
"lms": "LMS",
|
|
||||||
"dpmpp_2s_ancestral": "DPM++ 2S a",
|
|
||||||
"dpmpp_sde": "DPM++ SDE",
|
|
||||||
"dpmpp_sde_gpu": "DPM++ SDE",
|
|
||||||
"dpmpp_2m": "DPM++ 2M",
|
|
||||||
"dpmpp_2m_sde": "DPM++ 2M SDE",
|
|
||||||
"dpmpp_2m_sde_gpu": "DPM++ 2M SDE",
|
|
||||||
"ddim": "DDIM",
|
|
||||||
}
|
|
||||||
sampler_name = sampler_mapping.get(sampler, sampler)
|
|
||||||
|
|
||||||
if "scheduler" in metadata_dict:
|
|
||||||
scheduler = metadata_dict.get("scheduler")
|
|
||||||
scheduler_mapping = {
|
|
||||||
"normal": "Simple",
|
|
||||||
"karras": "Karras",
|
|
||||||
"exponential": "Exponential",
|
|
||||||
"sgm_uniform": "SGM Uniform",
|
|
||||||
"sgm_quadratic": "SGM Quadratic",
|
|
||||||
}
|
|
||||||
scheduler_name = scheduler_mapping.get(scheduler, scheduler)
|
|
||||||
|
|
||||||
# Add combined sampler and scheduler information
|
|
||||||
if sampler_name:
|
if sampler_name:
|
||||||
if scheduler_name:
|
if scheduler_name:
|
||||||
params.append(f"Sampler: {sampler_name} {scheduler_name}")
|
params.append(f"Sampler: {sampler_name} {scheduler_name}")
|
||||||
else:
|
else:
|
||||||
params.append(f"Sampler: {sampler_name}")
|
params.append(f"Sampler: {sampler_name}")
|
||||||
|
if cfg is not None:
|
||||||
|
params.append(f"CFG scale: {cfg}")
|
||||||
|
if seed is not None:
|
||||||
|
params.append(f"Seed: {seed}")
|
||||||
|
if size:
|
||||||
|
params.append(f"Size: {size}")
|
||||||
|
if clip_skip is not None:
|
||||||
|
try:
|
||||||
|
params.append(f"Clip skip: {abs(int(clip_skip))}")
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
pass
|
||||||
|
additional_data = metadata_dict.get("additional_data", "")
|
||||||
|
if additional_data:
|
||||||
|
params.append(additional_data)
|
||||||
|
if ckpt_hash:
|
||||||
|
params.append(f"Model hash: {ckpt_hash[:10].upper()}")
|
||||||
|
if ckpt_display_name:
|
||||||
|
params.append(f"Model: {ckpt_display_name}")
|
||||||
|
if hashes:
|
||||||
|
params.append(f"Hashes: {json.dumps(hashes, separators=(',', ':'))}")
|
||||||
|
params.append("Version: ComfyUI")
|
||||||
|
if civitai_resources:
|
||||||
|
params.append(
|
||||||
|
f"Civitai resources: {json.dumps(civitai_resources, separators=(',', ':'))}"
|
||||||
|
)
|
||||||
|
|
||||||
# CFG scale (Use guidance if available, otherwise fall back to cfg_scale or cfg)
|
lines.append(", ".join(params))
|
||||||
if "guidance" in metadata_dict:
|
return "\n".join(lines)
|
||||||
add_param_if_not_none(params, "CFG scale", metadata_dict.get("guidance"))
|
|
||||||
elif "cfg_scale" in metadata_dict:
|
|
||||||
add_param_if_not_none(params, "CFG scale", metadata_dict.get("cfg_scale"))
|
|
||||||
elif "cfg" in metadata_dict:
|
|
||||||
add_param_if_not_none(params, "CFG scale", metadata_dict.get("cfg"))
|
|
||||||
|
|
||||||
# Seed
|
|
||||||
if "seed" in metadata_dict:
|
|
||||||
add_param_if_not_none(params, "Seed", metadata_dict.get("seed"))
|
|
||||||
|
|
||||||
# Size
|
|
||||||
if "size" in metadata_dict:
|
|
||||||
add_param_if_not_none(params, "Size", metadata_dict.get("size"))
|
|
||||||
|
|
||||||
# Model info
|
|
||||||
if "checkpoint" in metadata_dict:
|
|
||||||
# Ensure checkpoint is a string before processing
|
|
||||||
checkpoint = metadata_dict.get("checkpoint")
|
|
||||||
if checkpoint is not None:
|
|
||||||
# Get model hash
|
|
||||||
model_hash = self.get_checkpoint_hash(checkpoint)
|
|
||||||
|
|
||||||
# Extract basename without path
|
|
||||||
checkpoint_name = os.path.basename(checkpoint)
|
|
||||||
# Remove extension if present
|
|
||||||
checkpoint_name = os.path.splitext(checkpoint_name)[0]
|
|
||||||
|
|
||||||
# Add model hash if available
|
|
||||||
if model_hash:
|
|
||||||
params.append(
|
|
||||||
f"Model hash: {model_hash[:10]}, Model: {checkpoint_name}"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
params.append(f"Model: {checkpoint_name}")
|
|
||||||
|
|
||||||
# Add LoRA hashes if available
|
|
||||||
if lora_hashes:
|
|
||||||
lora_hash_parts = []
|
|
||||||
for lora_name, hash_value in lora_hashes.items():
|
|
||||||
lora_hash_parts.append(f"{lora_name}: {hash_value[:10]}")
|
|
||||||
|
|
||||||
if lora_hash_parts:
|
|
||||||
params.append(f'Lora hashes: "{", ".join(lora_hash_parts)}"')
|
|
||||||
|
|
||||||
# Combine all parameters with commas
|
|
||||||
metadata_parts.append(", ".join(params))
|
|
||||||
|
|
||||||
# Join all parts with a new line
|
|
||||||
return "\n".join(metadata_parts)
|
|
||||||
|
|
||||||
# credit to nkchocoai
|
# credit to nkchocoai
|
||||||
# Add format_filename method to handle pattern substitution
|
# Add format_filename method to handle pattern substitution
|
||||||
@@ -573,6 +787,8 @@ class SaveImageLM:
|
|||||||
extra_pnginfo=None,
|
extra_pnginfo=None,
|
||||||
lossless_webp=True,
|
lossless_webp=True,
|
||||||
quality=100,
|
quality=100,
|
||||||
|
webp_method=6,
|
||||||
|
jpeg_subsampling=0,
|
||||||
embed_workflow=False,
|
embed_workflow=False,
|
||||||
save_with_metadata=True,
|
save_with_metadata=True,
|
||||||
add_counter_to_filename=True,
|
add_counter_to_filename=True,
|
||||||
@@ -627,15 +843,14 @@ class SaveImageLM:
|
|||||||
elif file_format == "jpeg":
|
elif file_format == "jpeg":
|
||||||
file = base_filename + ".jpg"
|
file = base_filename + ".jpg"
|
||||||
file_extension = ".jpg"
|
file_extension = ".jpg"
|
||||||
save_kwargs = {"quality": quality, "optimize": True}
|
save_kwargs = {"quality": quality, "optimize": True, "subsampling": jpeg_subsampling}
|
||||||
elif file_format == "webp":
|
elif file_format == "webp":
|
||||||
file = base_filename + ".webp"
|
file = base_filename + ".webp"
|
||||||
file_extension = ".webp"
|
file_extension = ".webp"
|
||||||
# Add optimization param to control performance
|
|
||||||
save_kwargs = {
|
save_kwargs = {
|
||||||
"quality": quality,
|
"quality": quality,
|
||||||
"lossless": lossless_webp,
|
"lossless": lossless_webp,
|
||||||
"method": 0,
|
"method": webp_method,
|
||||||
}
|
}
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported file format: {file_format}")
|
raise ValueError(f"Unsupported file format: {file_format}")
|
||||||
@@ -722,6 +937,8 @@ class SaveImageLM:
|
|||||||
extra_pnginfo=None,
|
extra_pnginfo=None,
|
||||||
lossless_webp=True,
|
lossless_webp=True,
|
||||||
quality=100,
|
quality=100,
|
||||||
|
webp_method=6,
|
||||||
|
jpeg_subsampling=0,
|
||||||
embed_workflow=False,
|
embed_workflow=False,
|
||||||
save_with_metadata=True,
|
save_with_metadata=True,
|
||||||
add_counter_to_filename=True,
|
add_counter_to_filename=True,
|
||||||
@@ -751,6 +968,8 @@ class SaveImageLM:
|
|||||||
extra_pnginfo,
|
extra_pnginfo,
|
||||||
lossless_webp,
|
lossless_webp,
|
||||||
quality,
|
quality,
|
||||||
|
webp_method,
|
||||||
|
jpeg_subsampling,
|
||||||
embed_workflow,
|
embed_workflow,
|
||||||
save_with_metadata,
|
save_with_metadata,
|
||||||
add_counter_to_filename,
|
add_counter_to_filename,
|
||||||
|
|||||||
@@ -36,6 +36,7 @@ any_type = AnyType("*")
|
|||||||
|
|
||||||
# Common methods extracted from lora_loader.py and lora_stacker.py
|
# Common methods extracted from lora_loader.py and lora_stacker.py
|
||||||
import os
|
import os
|
||||||
|
import re
|
||||||
import logging
|
import logging
|
||||||
import copy
|
import copy
|
||||||
import sys
|
import sys
|
||||||
@@ -69,6 +70,25 @@ def extract_lora_name(lora_path):
|
|||||||
return apply_lora_syntax_format(name_no_ext)
|
return apply_lora_syntax_format(name_no_ext)
|
||||||
|
|
||||||
|
|
||||||
|
def parse_lora_syntax(text: str) -> list[dict]:
|
||||||
|
"""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):
|
def get_loras_list(kwargs):
|
||||||
"""Helper to extract loras list from either old or new kwargs format"""
|
"""Helper to extract loras list from either old or new kwargs format"""
|
||||||
if "loras" not in kwargs:
|
if "loras" not in kwargs:
|
||||||
|
|||||||
@@ -0,0 +1,165 @@
|
|||||||
|
"""HTTP route handlers for agent skill endpoints.
|
||||||
|
|
||||||
|
These handlers expose the :class:`AgentService` via HTTP, allowing the
|
||||||
|
frontend to list available skills and execute them on selected models.
|
||||||
|
Progress is reported via WebSocket broadcast.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
from typing import Any, Dict
|
||||||
|
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
from ...services.agent import AgentService, AgentProgressReporter
|
||||||
|
from ...services.llm_service import LLMNotConfiguredError
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class AgentHandler:
|
||||||
|
"""HTTP handler for agent skill operations."""
|
||||||
|
|
||||||
|
def __init__(self, agent_service: AgentService | None = None) -> None:
|
||||||
|
self._agent_service = agent_service
|
||||||
|
|
||||||
|
async def _ensure_service(self) -> AgentService:
|
||||||
|
if self._agent_service is None:
|
||||||
|
self._agent_service = await AgentService.get_instance()
|
||||||
|
return self._agent_service
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# GET /api/lm/agent/skills
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def get_agent_skills(self, request: web.Request) -> web.Response:
|
||||||
|
"""Return a list of available agent skills."""
|
||||||
|
|
||||||
|
service = await self._ensure_service()
|
||||||
|
skills = await service.list_skills()
|
||||||
|
return web.json_response({"skills": skills})
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# POST /api/lm/agent/execute/{skill_name}
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def execute_agent_skill(self, request: web.Request) -> web.Response:
|
||||||
|
"""Execute an agent skill on the provided model paths.
|
||||||
|
|
||||||
|
Request body::
|
||||||
|
|
||||||
|
{"model_paths": ["/path/to/model1.safetensors", ...], "options": {}}
|
||||||
|
|
||||||
|
Returns immediately with a task ID. Execution runs in the
|
||||||
|
background; progress and completion are pushed via WebSocket
|
||||||
|
events of type ``agent_progress``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
skill_name = request.match_info.get("skill_name", "")
|
||||||
|
if not skill_name:
|
||||||
|
return web.json_response(
|
||||||
|
{"error": "Skill name is required"}, status=400
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
body = await request.json()
|
||||||
|
except Exception:
|
||||||
|
return web.json_response(
|
||||||
|
{"error": "Invalid JSON body"}, status=400
|
||||||
|
)
|
||||||
|
|
||||||
|
model_paths = body.get("model_paths", [])
|
||||||
|
if not model_paths or not isinstance(model_paths, list):
|
||||||
|
return web.json_response(
|
||||||
|
{"error": "model_paths must be a non-empty array"},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
service = await self._ensure_service()
|
||||||
|
|
||||||
|
# Validate LLM configuration early for skills that need it
|
||||||
|
# (fail fast rather than after starting background work)
|
||||||
|
try:
|
||||||
|
from ...services.llm_service import LLMService
|
||||||
|
|
||||||
|
llm = await LLMService.get_instance()
|
||||||
|
if not llm.is_configured():
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"error": "LLM provider is not configured. "
|
||||||
|
"Enable it in Settings → AI Provider.",
|
||||||
|
},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("Failed to check LLM configuration: %s", exc)
|
||||||
|
|
||||||
|
# Launch execution in the background
|
||||||
|
progress_reporter = AgentProgressReporter()
|
||||||
|
logger.info(
|
||||||
|
"LLM enrichment '%s' starting for %d model(s)",
|
||||||
|
skill_name, len(model_paths),
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _run() -> None:
|
||||||
|
try:
|
||||||
|
result = await service.execute_skill(
|
||||||
|
skill_name=skill_name,
|
||||||
|
input_data={"model_paths": model_paths},
|
||||||
|
progress_callback=progress_reporter,
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"LLM enrichment '%s' finished: success=%s, summary='%s', errors=%s",
|
||||||
|
skill_name, result.success, result.summary, result.errors,
|
||||||
|
)
|
||||||
|
except LLMNotConfiguredError as exc:
|
||||||
|
logger.warning("LLM enrichment '%s' not configured: %s", skill_name, exc)
|
||||||
|
await progress_reporter.on_progress(
|
||||||
|
{
|
||||||
|
"type": "agent_progress",
|
||||||
|
"skill": skill_name,
|
||||||
|
"status": "error",
|
||||||
|
"error": str(exc),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("LLM enrichment '%s' failed: %s", skill_name, exc, exc_info=True)
|
||||||
|
await progress_reporter.on_progress(
|
||||||
|
{
|
||||||
|
"type": "agent_progress",
|
||||||
|
"skill": skill_name,
|
||||||
|
"status": "error",
|
||||||
|
"error": str(exc),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
# Fire and forget — progress comes via WebSocket
|
||||||
|
asyncio.create_task(_run())
|
||||||
|
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"status": "started",
|
||||||
|
"skill": skill_name,
|
||||||
|
"model_count": len(model_paths),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# POST /api/lm/agent/cancel
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def cancel_agent_skill(self, request: web.Request) -> web.Response:
|
||||||
|
"""Cancel a running agent skill.
|
||||||
|
|
||||||
|
NOTE: Cancellation is a stub for now — the AgentService processes
|
||||||
|
models sequentially and does not yet support mid-execution
|
||||||
|
cancellation. This endpoint exists for API completeness.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# TODO: implement cooperative cancellation in AgentService
|
||||||
|
return web.json_response(
|
||||||
|
{"status": "acknowledged", "note": "Cancellation not yet implemented"},
|
||||||
|
status=200,
|
||||||
|
)
|
||||||
@@ -49,6 +49,14 @@ async def _get_hf_api_session() -> aiohttp.ClientSession:
|
|||||||
return _hf_api_session
|
return _hf_api_session
|
||||||
|
|
||||||
|
|
||||||
|
async def close_hf_api_session() -> None:
|
||||||
|
"""Close the shared HF API session, if it was ever created."""
|
||||||
|
global _hf_api_session
|
||||||
|
if _hf_api_session is not None and not _hf_api_session.closed:
|
||||||
|
await _hf_api_session.close()
|
||||||
|
_hf_api_session = None
|
||||||
|
|
||||||
|
|
||||||
def _infer_model_type(model_root: str) -> tuple[Any, str]:
|
def _infer_model_type(model_root: str) -> tuple[Any, str]:
|
||||||
"""Determine model class and scanner by matching ``model_root`` against the
|
"""Determine model class and scanner by matching ``model_root`` against the
|
||||||
configured root paths for each model type (from ``Config``).
|
configured root paths for each model type (from ``Config``).
|
||||||
@@ -114,8 +122,12 @@ async def _save_hf_metadata(dest_path: str, repo: str, model_root: str) -> None:
|
|||||||
metadata._unknown_fields["hf_url"] = hf_url
|
metadata._unknown_fields["hf_url"] = hf_url
|
||||||
metadata.from_civitai = False # HF models are not from CivitAI
|
metadata.from_civitai = False # HF models are not from CivitAI
|
||||||
|
|
||||||
|
metadata_dict = metadata.to_dict()
|
||||||
|
if "trainedWords" in metadata_dict and not metadata_dict["trainedWords"]:
|
||||||
|
del metadata_dict["trainedWords"]
|
||||||
|
|
||||||
# 3. Save metadata atomically
|
# 3. Save metadata atomically
|
||||||
await MetadataManager.save_metadata(dest_path, metadata)
|
await MetadataManager.save_metadata(dest_path, metadata_dict)
|
||||||
logger.info("Saved HF metadata (with hf_url) for %s", dest_path)
|
logger.info("Saved HF metadata (with hf_url) for %s", dest_path)
|
||||||
|
|
||||||
# 4. Determine relative folder path for cache
|
# 4. Determine relative folder path for cache
|
||||||
@@ -139,9 +151,117 @@ async def _save_hf_metadata(dest_path: str, repo: str, model_root: str) -> None:
|
|||||||
logger.warning("Failed to save HF metadata for %s: %s", dest_path, exc)
|
logger.warning("Failed to save HF metadata for %s: %s", dest_path, exc)
|
||||||
|
|
||||||
|
|
||||||
|
def _find_matching_root(dest_dir: str) -> str | None:
|
||||||
|
"""Walk up *dest_dir* to find which configured scanner root it belongs to."""
|
||||||
|
norm = os.path.normpath(dest_dir).replace(os.sep, "/")
|
||||||
|
all_roots = []
|
||||||
|
for root_list in (
|
||||||
|
config.loras_roots or [],
|
||||||
|
config.extra_loras_roots or [],
|
||||||
|
config.checkpoints_roots or [],
|
||||||
|
config.extra_checkpoints_roots or [],
|
||||||
|
config.unet_roots or [],
|
||||||
|
config.extra_unet_roots or [],
|
||||||
|
config.embeddings_roots or [],
|
||||||
|
config.extra_embeddings_roots or [],
|
||||||
|
):
|
||||||
|
all_roots.extend([os.path.normpath(p).replace(os.sep, "/") for p in root_list])
|
||||||
|
# Find the longest matching prefix
|
||||||
|
match: str | None = None
|
||||||
|
for root in all_roots:
|
||||||
|
if norm.startswith(root):
|
||||||
|
if match is None or len(root) > len(match):
|
||||||
|
match = root
|
||||||
|
return match
|
||||||
|
|
||||||
|
|
||||||
|
async def _add_to_scanner_cache(dest_path: str, metadata: dict[str, Any]) -> None:
|
||||||
|
model_dir = os.path.dirname(dest_path)
|
||||||
|
model_root = _find_matching_root(model_dir)
|
||||||
|
if not model_root:
|
||||||
|
raise ValueError(f"File path {dest_path} is not within any configured scanner root")
|
||||||
|
scanner_getter_name = _infer_model_type(model_root)[1]
|
||||||
|
scanner_getter = getattr(ServiceRegistry, scanner_getter_name, None)
|
||||||
|
if scanner_getter is None:
|
||||||
|
raise RuntimeError(f"Scanner getter '{scanner_getter_name}' not found in ServiceRegistry")
|
||||||
|
scanner = await scanner_getter()
|
||||||
|
if scanner is None:
|
||||||
|
raise RuntimeError(f"Scanner '{scanner_getter_name}' returned None")
|
||||||
|
await scanner.update_single_model_cache(dest_path, dest_path, metadata)
|
||||||
|
|
||||||
|
|
||||||
class HfHandler:
|
class HfHandler:
|
||||||
"""Handle Hugging Face model browsing and download."""
|
"""Handle Hugging Face model browsing and download."""
|
||||||
|
|
||||||
|
async def set_hf_url(self, request: web.Request) -> web.Response:
|
||||||
|
try:
|
||||||
|
payload: dict[str, Any] = await request.json()
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
return web.json_response({"success": False, "error": "Invalid JSON"}, status=400)
|
||||||
|
|
||||||
|
file_path = (payload.get("file_path") or "").strip()
|
||||||
|
hf_url = (payload.get("hf_url") or "").strip()
|
||||||
|
|
||||||
|
if not file_path or not hf_url:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Missing required fields: 'file_path' and 'hf_url'"},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
m = re.match(r"^https?://huggingface\.co/([^/]+/[^/]+)/?$", hf_url)
|
||||||
|
if not m:
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": False,
|
||||||
|
"error": "Invalid HuggingFace URL. Expected format: https://huggingface.co/user/repo",
|
||||||
|
},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not os.path.isfile(file_path):
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": f"File not found: {file_path}"},
|
||||||
|
status=404,
|
||||||
|
)
|
||||||
|
|
||||||
|
model_root = _find_matching_root(os.path.dirname(file_path))
|
||||||
|
if not model_root:
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": False,
|
||||||
|
"error": "File is not within any configured model directory. Cannot link to HuggingFace.",
|
||||||
|
},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
existing = await MetadataManager.load_metadata_payload(file_path)
|
||||||
|
if existing.get("hf_url") == hf_url:
|
||||||
|
return web.json_response({
|
||||||
|
"success": True,
|
||||||
|
"message": "hf_url already set",
|
||||||
|
"hf_url": hf_url,
|
||||||
|
})
|
||||||
|
|
||||||
|
existing["hf_url"] = hf_url
|
||||||
|
existing["from_civitai"] = False
|
||||||
|
await MetadataManager.save_metadata(file_path, existing)
|
||||||
|
|
||||||
|
await _add_to_scanner_cache(file_path, existing)
|
||||||
|
|
||||||
|
logger.info("Set hf_url=%s for %s", hf_url, file_path)
|
||||||
|
return web.json_response({
|
||||||
|
"success": True,
|
||||||
|
"message": f"hf_url set to {hf_url}",
|
||||||
|
"hf_url": hf_url,
|
||||||
|
})
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("Failed to set hf_url for %s: %s", file_path, exc)
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": str(exc)},
|
||||||
|
status=500,
|
||||||
|
)
|
||||||
|
|
||||||
async def get_hf_repo_files(self, request: web.Request) -> web.Response:
|
async def get_hf_repo_files(self, request: web.Request) -> web.Response:
|
||||||
"""List model-weight files from a HF repo with real file sizes.
|
"""List model-weight files from a HF repo with real file sizes.
|
||||||
|
|
||||||
@@ -243,8 +363,8 @@ class HfHandler:
|
|||||||
if ".." in (author, repo_name) or "." in (author, repo_name):
|
if ".." in (author, repo_name) or "." in (author, repo_name):
|
||||||
return web.json_response({"error": f"Invalid repo format: {repo}"}, status=400)
|
return web.json_response({"error": f"Invalid repo format: {repo}"}, status=400)
|
||||||
|
|
||||||
# Validate filename — must not contain path separators or ..
|
# Validate filename — must not contain path traversal
|
||||||
if "/" in filename or "\\" in filename or ".." in filename:
|
if ".." in filename:
|
||||||
return web.json_response({"error": "Invalid filename"}, status=400)
|
return web.json_response({"error": "Invalid filename"}, status=400)
|
||||||
|
|
||||||
# Validate relative_path — must not be absolute or escape base directory
|
# Validate relative_path — must not be absolute or escape base directory
|
||||||
@@ -254,35 +374,17 @@ class HfHandler:
|
|||||||
if ".." in relative_path.split("/") or "\\" in relative_path:
|
if ".." in relative_path.split("/") or "\\" in relative_path:
|
||||||
return web.json_response({"error": "Invalid relative_path"}, status=400)
|
return web.json_response({"error": "Invalid relative_path"}, status=400)
|
||||||
|
|
||||||
# Validate model_root — must not contain path traversal
|
# Use model_root directly as the base directory — same approach as
|
||||||
if not os.path.isabs(model_root):
|
# CivitAI's download path (download_manager.py). No realpath, no
|
||||||
# For relative model_root, check it doesn't escape
|
# allowed-roots validation, no path-traversal check; those are
|
||||||
resolved_model_root = os.path.realpath(
|
# unnecessary when the frontend sends the path from its own dropdown
|
||||||
os.path.join(os.getcwd(), "models", model_root)
|
# (populated from scanner roots). Using the "business path" directly
|
||||||
)
|
# keeps dest_path consistent with scanner roots so that later folder
|
||||||
|
# derivation (in _save_hf_metadata) works correctly.
|
||||||
|
if os.path.isabs(model_root):
|
||||||
|
base_dir = os.path.normpath(model_root)
|
||||||
else:
|
else:
|
||||||
resolved_model_root = os.path.realpath(model_root)
|
base_dir = os.path.normpath(os.path.join(os.getcwd(), "models", model_root))
|
||||||
|
|
||||||
# Verify model_root is within a configured scanner root
|
|
||||||
allowed_roots = set()
|
|
||||||
for root_list in (
|
|
||||||
config.loras_roots or [],
|
|
||||||
config.extra_loras_roots or [],
|
|
||||||
config.checkpoints_roots or [],
|
|
||||||
config.extra_checkpoints_roots or [],
|
|
||||||
config.unet_roots or [],
|
|
||||||
config.extra_unet_roots or [],
|
|
||||||
config.embeddings_roots or [],
|
|
||||||
config.extra_embeddings_roots or [],
|
|
||||||
):
|
|
||||||
for r in root_list:
|
|
||||||
allowed_roots.add(os.path.realpath(r))
|
|
||||||
|
|
||||||
if not any(resolved_model_root == root or resolved_model_root.startswith(root + os.sep) for root in allowed_roots):
|
|
||||||
logger.warning("Invalid model_root rejected: %s", model_root)
|
|
||||||
return web.json_response({"error": f"Invalid model_root: {model_root}"}, status=400)
|
|
||||||
|
|
||||||
base_dir = resolved_model_root
|
|
||||||
|
|
||||||
if use_default_paths:
|
if use_default_paths:
|
||||||
target_dir = os.path.join(base_dir, "huggingface", author, repo_name)
|
target_dir = os.path.join(base_dir, "huggingface", author, repo_name)
|
||||||
@@ -291,15 +393,12 @@ class HfHandler:
|
|||||||
else:
|
else:
|
||||||
target_dir = base_dir
|
target_dir = base_dir
|
||||||
|
|
||||||
os.makedirs(target_dir, exist_ok=True)
|
# Strip HF repo subdirectory — "diffusion_models/xxx.safetensors"
|
||||||
dest_path = os.path.join(target_dir, filename)
|
# is an HF repo convention, not meaningful for local storage.
|
||||||
|
file_base = os.path.basename(filename)
|
||||||
|
|
||||||
# Resolve symlinks and check for path traversal escape
|
os.makedirs(target_dir, exist_ok=True)
|
||||||
real_dest = os.path.realpath(dest_path)
|
dest_path = os.path.join(target_dir, file_base)
|
||||||
real_base = os.path.realpath(target_dir)
|
|
||||||
if not real_dest.startswith(real_base + os.sep):
|
|
||||||
logger.warning("Path traversal blocked: %s -> %s", dest_path, real_dest)
|
|
||||||
return web.json_response({"error": "Path traversal detected"}, status=400)
|
|
||||||
|
|
||||||
# Check if already exists (simple skip)
|
# Check if already exists (simple skip)
|
||||||
if os.path.exists(dest_path) and os.path.getsize(dest_path) > 0:
|
if os.path.exists(dest_path) and os.path.getsize(dest_path) > 0:
|
||||||
|
|||||||
@@ -38,6 +38,12 @@ from ...services.settings_manager import get_settings_manager
|
|||||||
from ...services.websocket_manager import ws_manager
|
from ...services.websocket_manager import ws_manager
|
||||||
from ...services.downloader import get_downloader
|
from ...services.downloader import get_downloader
|
||||||
from ...services.errors import ResourceNotFoundError
|
from ...services.errors import ResourceNotFoundError
|
||||||
|
from ...services.llm_service import (
|
||||||
|
PROVIDER_PRESETS,
|
||||||
|
fetch_ollama_models,
|
||||||
|
get_all_provider_models,
|
||||||
|
get_provider_model_ids,
|
||||||
|
)
|
||||||
from ...services.cache_health_monitor import CacheHealthMonitor, CacheHealthStatus
|
from ...services.cache_health_monitor import CacheHealthMonitor, CacheHealthStatus
|
||||||
from ...utils.models import BaseModelMetadata
|
from ...utils.models import BaseModelMetadata
|
||||||
from ...utils.constants import (
|
from ...utils.constants import (
|
||||||
@@ -49,6 +55,7 @@ from ...utils.constants import (
|
|||||||
VALID_LORA_TYPES,
|
VALID_LORA_TYPES,
|
||||||
)
|
)
|
||||||
from .hf_handlers import HfHandler
|
from .hf_handlers import HfHandler
|
||||||
|
from .agent_handlers import AgentHandler
|
||||||
from ...utils.civitai_utils import rewrite_preview_url
|
from ...utils.civitai_utils import rewrite_preview_url
|
||||||
from ...utils.example_images_paths import (
|
from ...utils.example_images_paths import (
|
||||||
find_non_compliant_items_in_example_images_root,
|
find_non_compliant_items_in_example_images_root,
|
||||||
@@ -566,12 +573,18 @@ class NodeRegistry:
|
|||||||
tab_nodes[nd["unique_id"]] = nd
|
tab_nodes[nd["unique_id"]] = nd
|
||||||
|
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
|
prev_count = len(self._tab_nodes.get(sid, {}))
|
||||||
self._tab_nodes[sid] = tab_nodes
|
self._tab_nodes[sid] = tab_nodes
|
||||||
self._waiting_clients.discard(sid)
|
self._waiting_clients.discard(sid)
|
||||||
if not self._waiting_clients:
|
if not self._waiting_clients:
|
||||||
self._ready.set()
|
self._ready.set()
|
||||||
|
total_tabs = len(self._tab_nodes)
|
||||||
|
|
||||||
logger.debug("Registered %s nodes from client %s", len(nodes), sid)
|
if len(nodes) != prev_count or len(nodes) > 0:
|
||||||
|
logger.debug(
|
||||||
|
"[LM:Registry] stored %s nodes (was %s) for client %s (total tabs: %s)",
|
||||||
|
len(nodes), prev_count, sid, total_tabs,
|
||||||
|
)
|
||||||
|
|
||||||
def prepare_for_refresh(self, active_sids: list[str]) -> None:
|
def prepare_for_refresh(self, active_sids: list[str]) -> None:
|
||||||
"""Set the list of client IDs we expect to hear from during the next refresh cycle."""
|
"""Set the list of client IDs we expect to hear from during the next refresh cycle."""
|
||||||
@@ -594,10 +607,17 @@ class NodeRegistry:
|
|||||||
longer connected."""
|
longer connected."""
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
# Garbage-collect stale entries (disconnected tabs)
|
# Garbage-collect stale entries (disconnected tabs)
|
||||||
|
stale_sids = []
|
||||||
if active_sids is not None:
|
if active_sids is not None:
|
||||||
for sid in list(self._tab_nodes):
|
for sid in list(self._tab_nodes):
|
||||||
if sid not in active_sids:
|
if sid not in active_sids:
|
||||||
|
stale_sids.append(sid)
|
||||||
del self._tab_nodes[sid]
|
del self._tab_nodes[sid]
|
||||||
|
if stale_sids:
|
||||||
|
logger.debug(
|
||||||
|
"[LM:Registry] GC pruned %s disconnected tabs: %s",
|
||||||
|
len(stale_sids), stale_sids,
|
||||||
|
)
|
||||||
|
|
||||||
merged: dict[str, dict] = {}
|
merged: dict[str, dict] = {}
|
||||||
tab_info: dict[str, dict] = {}
|
tab_info: dict[str, dict] = {}
|
||||||
@@ -1399,8 +1419,9 @@ class SettingsHandler:
|
|||||||
"libraries",
|
"libraries",
|
||||||
"active_library",
|
"active_library",
|
||||||
# Sensitive — never expose the actual value to the frontend;
|
# Sensitive — never expose the actual value to the frontend;
|
||||||
# frontend receives a boolean instead (civitai_api_key_set).
|
# frontend receives a boolean instead (*_set).
|
||||||
"civitai_api_key",
|
"civitai_api_key",
|
||||||
|
"llm_api_key",
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1458,6 +1479,8 @@ class SettingsHandler:
|
|||||||
# Sensitive fields: only expose a boolean indicating whether set
|
# Sensitive fields: only expose a boolean indicating whether set
|
||||||
raw_key = self._settings.get("civitai_api_key")
|
raw_key = self._settings.get("civitai_api_key")
|
||||||
response_data["civitai_api_key_set"] = bool(raw_key)
|
response_data["civitai_api_key_set"] = bool(raw_key)
|
||||||
|
raw_llm_key = self._settings.get("llm_api_key")
|
||||||
|
response_data["llm_api_key_set"] = bool(raw_llm_key)
|
||||||
settings_file = getattr(self._settings, "settings_file", None)
|
settings_file = getattr(self._settings, "settings_file", None)
|
||||||
if settings_file:
|
if settings_file:
|
||||||
response_data["settings_file"] = settings_file
|
response_data["settings_file"] = settings_file
|
||||||
@@ -1539,6 +1562,11 @@ class SettingsHandler:
|
|||||||
{"success": False, "error": validation_error}
|
{"success": False, "error": validation_error}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if key == "update_channel" and value not in ("release", "nightly"):
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "update_channel must be 'release' or 'nightly'"}
|
||||||
|
)
|
||||||
|
|
||||||
if value == "__DELETE__" and key in (
|
if value == "__DELETE__" and key in (
|
||||||
"proxy_username",
|
"proxy_username",
|
||||||
"proxy_password",
|
"proxy_password",
|
||||||
@@ -1547,7 +1575,11 @@ class SettingsHandler:
|
|||||||
else:
|
else:
|
||||||
self._settings.set(key, value)
|
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()
|
await self._metadata_provider_updater()
|
||||||
|
|
||||||
if key in self._PROXY_KEYS:
|
if key in self._PROXY_KEYS:
|
||||||
@@ -1562,6 +1594,42 @@ class SettingsHandler:
|
|||||||
logger.error("Error updating settings: %s", exc, exc_info=True)
|
logger.error("Error updating settings: %s", exc, exc_info=True)
|
||||||
return web.Response(status=500, text=str(exc))
|
return web.Response(status=500, text=str(exc))
|
||||||
|
|
||||||
|
async def get_llm_models(self, request: web.Request) -> web.Response:
|
||||||
|
"""Return the model list for a provider.
|
||||||
|
|
||||||
|
For ``ollama`` the list is fetched live from the local Ollama API
|
||||||
|
(only models actually pulled locally are shown). For all other
|
||||||
|
providers the opencode model catalog is used.
|
||||||
|
|
||||||
|
Query parameters:
|
||||||
|
provider (required): Internal provider id (``openai``, ``ollama``, etc.).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``{"success": true, "models": ["gpt-4o", ...]}``.
|
||||||
|
"""
|
||||||
|
provider_id = request.query.get("provider", "").strip()
|
||||||
|
if not provider_id:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "provider query parameter is required", "models": []},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
if provider_id == "ollama":
|
||||||
|
api_base = request.query.get("api_base", "").strip() or self._settings.get("llm_api_base", "")
|
||||||
|
if not api_base:
|
||||||
|
api_base = "http://localhost:11434/v1"
|
||||||
|
models = await fetch_ollama_models(api_base)
|
||||||
|
else:
|
||||||
|
models = await get_provider_model_ids(provider_id)
|
||||||
|
return web.json_response({"success": True, "models": models})
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("get_llm_models failed for %s: %s", provider_id, exc)
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": str(exc), "models": []},
|
||||||
|
status=500,
|
||||||
|
)
|
||||||
|
|
||||||
def _validate_example_images_path(self, folder_path: str) -> str | None:
|
def _validate_example_images_path(self, folder_path: str) -> str | None:
|
||||||
if not os.path.exists(folder_path):
|
if not os.path.exists(folder_path):
|
||||||
return f"Path does not exist: {folder_path}"
|
return f"Path does not exist: {folder_path}"
|
||||||
@@ -1584,6 +1652,20 @@ class SettingsHandler:
|
|||||||
def _is_dedicated_example_images_folder(self, folder_path: str) -> bool:
|
def _is_dedicated_example_images_folder(self, folder_path: str) -> bool:
|
||||||
return is_valid_example_images_root(folder_path)
|
return is_valid_example_images_root(folder_path)
|
||||||
|
|
||||||
|
async def get_provider_models(self, request: web.Request) -> web.Response:
|
||||||
|
"""Return the model catalog for all preset providers.
|
||||||
|
|
||||||
|
This endpoint is called asynchronously by the settings UI so that
|
||||||
|
page rendering never blocks on the remote model catalog fetch.
|
||||||
|
"""
|
||||||
|
catalog_provider_ids = [p for p in PROVIDER_PRESETS if p != "custom"]
|
||||||
|
try:
|
||||||
|
provider_models = await get_all_provider_models(catalog_provider_ids)
|
||||||
|
return web.json_response({"success": True, "models": provider_models})
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Failed to fetch provider models: %s", exc)
|
||||||
|
return web.json_response({"success": False, "models": {}, "error": str(exc)})
|
||||||
|
|
||||||
|
|
||||||
class UsageStatsHandler:
|
class UsageStatsHandler:
|
||||||
def __init__(self, usage_stats_factory: UsageStatsFactory = UsageStats) -> None:
|
def __init__(self, usage_stats_factory: UsageStatsFactory = UsageStats) -> None:
|
||||||
@@ -1711,6 +1793,124 @@ class LoraCodeHandler:
|
|||||||
logger.error("Failed to update lora code: %s", exc, exc_info=True)
|
logger.error("Failed to update lora code: %s", exc, exc_info=True)
|
||||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
async def 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:
|
class TrainedWordsHandler:
|
||||||
async def get_trained_words(self, request: web.Request) -> web.Response:
|
async def get_trained_words(self, request: web.Request) -> web.Response:
|
||||||
@@ -3056,6 +3256,8 @@ class NodeRegistryHandler:
|
|||||||
self._node_registry = node_registry
|
self._node_registry = node_registry
|
||||||
self._prompt_server = prompt_server
|
self._prompt_server = prompt_server
|
||||||
self._standalone_mode = standalone_mode
|
self._standalone_mode = standalone_mode
|
||||||
|
self._refresh_lock = asyncio.Lock()
|
||||||
|
self._last_slow_path_ts: float = 0.0
|
||||||
|
|
||||||
async def register_nodes(self, request: web.Request) -> web.Response:
|
async def register_nodes(self, request: web.Request) -> web.Response:
|
||||||
try:
|
try:
|
||||||
@@ -3102,7 +3304,12 @@ class NodeRegistryHandler:
|
|||||||
)
|
)
|
||||||
graph_name = node.get("graph_name")
|
graph_name = node.get("graph_name")
|
||||||
try:
|
try:
|
||||||
node["node_id"] = int(node_id)
|
# Handle compound node IDs from expanded group subgraphs,
|
||||||
|
# e.g. "252:0" → 0 (parent scope is already in graph_id)
|
||||||
|
if isinstance(node_id, str) and ":" in node_id:
|
||||||
|
node["node_id"] = int(node_id.rsplit(":", 1)[-1])
|
||||||
|
else:
|
||||||
|
node["node_id"] = int(node_id)
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{
|
{
|
||||||
@@ -3143,42 +3350,101 @@ class NodeRegistryHandler:
|
|||||||
status=503,
|
status=503,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Snapshot of currently-connected ComfyUI tabs
|
|
||||||
active_sids = list(self._prompt_server.instance.sockets.keys())
|
|
||||||
self._node_registry.prepare_for_refresh(active_sids)
|
|
||||||
|
|
||||||
try:
|
|
||||||
self._prompt_server.instance.send_sync("lora_registry_refresh", {})
|
|
||||||
logger.debug(
|
|
||||||
"Sent registry refresh request (expecting %s clients)", len(active_sids)
|
|
||||||
)
|
|
||||||
except Exception as exc:
|
|
||||||
logger.error("Failed to send registry refresh message: %s", exc)
|
|
||||||
return web.json_response(
|
|
||||||
{
|
|
||||||
"success": False,
|
|
||||||
"error": "Communication Error",
|
|
||||||
"message": f"Failed to communicate with ComfyUI frontend: {exc}",
|
|
||||||
},
|
|
||||||
status=500,
|
|
||||||
)
|
|
||||||
|
|
||||||
if not await self._node_registry.wait_for_all(timeout=2.0):
|
|
||||||
logger.warning(
|
|
||||||
"Registry refresh timeout after 2s (%s/%s clients responded)",
|
|
||||||
len(active_sids) - self._node_registry.pending_client_count,
|
|
||||||
len(active_sids),
|
|
||||||
)
|
|
||||||
|
|
||||||
# Re-read current sockets after the wait: a tab may have connected
|
|
||||||
# while we were waiting, and we don't want to garbage-collect it.
|
|
||||||
current_sids = set(self._prompt_server.instance.sockets.keys())
|
current_sids = set(self._prompt_server.instance.sockets.keys())
|
||||||
|
|
||||||
|
# Fast path: if the frontend has already pushed node data (via
|
||||||
|
# afterConfigureGraph / graphChanged hooks), return it immediately
|
||||||
|
# without triggering a WebSocket round-trip.
|
||||||
registry_info = await self._node_registry.get_merged_registry(
|
registry_info = await self._node_registry.get_merged_registry(
|
||||||
active_sids=current_sids
|
active_sids=current_sids
|
||||||
)
|
)
|
||||||
|
if registry_info["tab_count"] > 0:
|
||||||
|
logger.debug(
|
||||||
|
"[LM:Registry] fast path: %s nodes across %s tabs %s",
|
||||||
|
registry_info["node_count"],
|
||||||
|
registry_info["tab_count"],
|
||||||
|
dict(registry_info.get("tabs", {})),
|
||||||
|
)
|
||||||
|
return web.json_response({"success": True, "data": registry_info})
|
||||||
|
|
||||||
|
# Slow path: registry is empty — trigger refresh via WebSocket.
|
||||||
|
# Serialize with an async lock so concurrent callers don't all
|
||||||
|
# trigger separate WS refresh cycles. The second caller will
|
||||||
|
# re-check the fast path and (usually) find populated data.
|
||||||
|
async with self._refresh_lock:
|
||||||
|
# Re-check after acquiring the lock — another concurrent call
|
||||||
|
# may have populated the cache while we were waiting.
|
||||||
|
registry_info = await self._node_registry.get_merged_registry(
|
||||||
|
active_sids=current_sids
|
||||||
|
)
|
||||||
|
if registry_info["tab_count"] > 0:
|
||||||
|
logger.debug(
|
||||||
|
"[LM:Registry] fast path after lock wait: %s nodes across %s tabs",
|
||||||
|
registry_info["node_count"],
|
||||||
|
registry_info["tab_count"],
|
||||||
|
)
|
||||||
|
return web.json_response({"success": True, "data": registry_info})
|
||||||
|
|
||||||
|
# Cooldown: if the slow path ran recently (< 2 s) and
|
||||||
|
# returned empty, skip another WS round-trip.
|
||||||
|
elapsed = time.monotonic() - self._last_slow_path_ts
|
||||||
|
if elapsed < 2.0:
|
||||||
|
logger.debug(
|
||||||
|
"[LM:Registry] slow path cooldown (%.1fs since last refresh), returning empty",
|
||||||
|
elapsed,
|
||||||
|
)
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": False,
|
||||||
|
"error": "Empty Registry",
|
||||||
|
"message": "No workflow nodes found — ensure ComfyUI is open and the extension is loaded.",
|
||||||
|
},
|
||||||
|
status=408,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.debug(
|
||||||
|
"[LM:Registry] slow path: cache empty, triggering WS refresh (%s connected tabs: %s)",
|
||||||
|
len(current_sids), list(current_sids)[:5],
|
||||||
|
)
|
||||||
|
active_sids = list(current_sids)
|
||||||
|
self._node_registry.prepare_for_refresh(active_sids)
|
||||||
|
|
||||||
|
try:
|
||||||
|
self._prompt_server.instance.send_sync("lora_registry_refresh", {})
|
||||||
|
logger.debug(
|
||||||
|
"Sent registry refresh request (expecting %s clients)", len(active_sids)
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("Failed to send registry refresh message: %s", exc)
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": False,
|
||||||
|
"error": "Communication Error",
|
||||||
|
"message": f"Failed to communicate with ComfyUI frontend: {exc}",
|
||||||
|
},
|
||||||
|
status=500,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not await self._node_registry.wait_for_all(timeout=0.5):
|
||||||
|
logger.warning(
|
||||||
|
"Registry refresh timeout after 0.5s (%s/%s clients responded)",
|
||||||
|
len(active_sids) - self._node_registry.pending_client_count,
|
||||||
|
len(active_sids),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Re-read current sockets after the wait: a tab may have connected
|
||||||
|
# while we were waiting, and we don't want to garbage-collect it.
|
||||||
|
current_sids = set(self._prompt_server.instance.sockets.keys())
|
||||||
|
registry_info = await self._node_registry.get_merged_registry(
|
||||||
|
active_sids=current_sids
|
||||||
|
)
|
||||||
|
self._last_slow_path_ts = time.monotonic()
|
||||||
|
|
||||||
if registry_info["node_count"] == 0:
|
if registry_info["node_count"] == 0:
|
||||||
logger.warning("No nodes registered after refresh")
|
logger.debug(
|
||||||
|
"[LM:Registry] refresh OK — %s connected tab(s) but 0 compatible nodes found",
|
||||||
|
registry_info["tab_count"],
|
||||||
|
)
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{
|
{
|
||||||
"success": False,
|
"success": False,
|
||||||
@@ -3214,7 +3480,7 @@ class NodeRegistryHandler:
|
|||||||
status=400,
|
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(
|
return web.json_response(
|
||||||
{"success": False, "error": "Missing value parameter"}, status=400
|
{"success": False, "error": "Missing value parameter"}, status=400
|
||||||
)
|
)
|
||||||
@@ -3292,6 +3558,130 @@ class NodeRegistryHandler:
|
|||||||
logger.error("Failed to update node widget: %s", exc, exc_info=True)
|
logger.error("Failed to update node widget: %s", exc, exc_info=True)
|
||||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
async def 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 = {
|
||||||
|
"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:
|
class MiscHandlerSet:
|
||||||
"""Aggregate handlers into a lookup compatible with the registrar."""
|
"""Aggregate handlers into a lookup compatible with the registrar."""
|
||||||
@@ -3317,6 +3707,7 @@ class MiscHandlerSet:
|
|||||||
example_workflows: ExampleWorkflowsHandler,
|
example_workflows: ExampleWorkflowsHandler,
|
||||||
base_model: BaseModelHandlerSet,
|
base_model: BaseModelHandlerSet,
|
||||||
hf_handler: HfHandler | None = None,
|
hf_handler: HfHandler | None = None,
|
||||||
|
agent_handler: AgentHandler | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.health = health
|
self.health = health
|
||||||
self.settings = settings
|
self.settings = settings
|
||||||
@@ -3336,6 +3727,7 @@ class MiscHandlerSet:
|
|||||||
self.example_workflows = example_workflows
|
self.example_workflows = example_workflows
|
||||||
self.base_model = base_model
|
self.base_model = base_model
|
||||||
self.hf_handler = hf_handler
|
self.hf_handler = hf_handler
|
||||||
|
self.agent_handler = agent_handler
|
||||||
|
|
||||||
def to_route_mapping(
|
def to_route_mapping(
|
||||||
self,
|
self,
|
||||||
@@ -3351,13 +3743,17 @@ class MiscHandlerSet:
|
|||||||
"get_priority_tags": self.settings.get_priority_tags,
|
"get_priority_tags": self.settings.get_priority_tags,
|
||||||
"get_settings_libraries": self.settings.get_libraries,
|
"get_settings_libraries": self.settings.get_libraries,
|
||||||
"activate_library": self.settings.activate_library,
|
"activate_library": self.settings.activate_library,
|
||||||
|
"get_llm_models": self.settings.get_llm_models,
|
||||||
|
"get_provider_models": self.settings.get_provider_models,
|
||||||
"update_usage_stats": self.usage_stats.update_usage_stats,
|
"update_usage_stats": self.usage_stats.update_usage_stats,
|
||||||
"get_usage_stats": self.usage_stats.get_usage_stats,
|
"get_usage_stats": self.usage_stats.get_usage_stats,
|
||||||
"update_lora_code": self.lora_code.update_lora_code,
|
"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_trained_words": self.trained_words.get_trained_words,
|
||||||
"get_model_example_files": self.model_examples.get_model_example_files,
|
"get_model_example_files": self.model_examples.get_model_example_files,
|
||||||
"register_nodes": self.node_registry.register_nodes,
|
"register_nodes": self.node_registry.register_nodes,
|
||||||
"update_node_widget": self.node_registry.update_node_widget,
|
"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,
|
"get_registry": self.node_registry.get_registry,
|
||||||
"check_model_exists": self.model_library.check_model_exists,
|
"check_model_exists": self.model_library.check_model_exists,
|
||||||
"check_models_exist": self.model_library.check_models_exist,
|
"check_models_exist": self.model_library.check_models_exist,
|
||||||
@@ -3384,6 +3780,11 @@ class MiscHandlerSet:
|
|||||||
# Hugging Face handlers
|
# Hugging Face handlers
|
||||||
"get_hf_repo_files": self.hf_handler.get_hf_repo_files,
|
"get_hf_repo_files": self.hf_handler.get_hf_repo_files,
|
||||||
"download_hf_model": self.hf_handler.download_hf_model,
|
"download_hf_model": self.hf_handler.download_hf_model,
|
||||||
|
"set_hf_url": self.hf_handler.set_hf_url,
|
||||||
|
# Agent skill handlers
|
||||||
|
"get_agent_skills": self.agent_handler.get_agent_skills,
|
||||||
|
"execute_agent_skill": self.agent_handler.execute_agent_skill,
|
||||||
|
"cancel_agent_skill": self.agent_handler.cancel_agent_skill,
|
||||||
# Base model handlers
|
# Base model handlers
|
||||||
"get_base_models": self.base_model.get_base_models,
|
"get_base_models": self.base_model.get_base_models,
|
||||||
"refresh_base_models": self.base_model.refresh_base_models,
|
"refresh_base_models": self.base_model.refresh_base_models,
|
||||||
|
|||||||
@@ -154,6 +154,14 @@ class ModelPageView:
|
|||||||
)
|
)
|
||||||
self._template_env._i18n_filter_added = True # type: ignore[attr-defined]
|
self._template_env._i18n_filter_added = True # type: ignore[attr-defined]
|
||||||
|
|
||||||
|
from ...services.llm_service import PROVIDER_PRESETS
|
||||||
|
|
||||||
|
# Provider presets are embedded directly (local, no await needed).
|
||||||
|
# Provider model catalogs are fetched asynchronously by the
|
||||||
|
# frontend via GET /api/lm/llm/provider-models so page rendering
|
||||||
|
# never blocks on the remote model catalog (which can take up to
|
||||||
|
# 30s on cold cache).
|
||||||
|
|
||||||
template_context = {
|
template_context = {
|
||||||
"is_initializing": is_initializing,
|
"is_initializing": is_initializing,
|
||||||
"settings": self._settings,
|
"settings": self._settings,
|
||||||
@@ -161,6 +169,8 @@ class ModelPageView:
|
|||||||
"folders": [],
|
"folders": [],
|
||||||
"t": self._server_i18n.get_translation,
|
"t": self._server_i18n.get_translation,
|
||||||
"version": self._get_app_version(),
|
"version": self._get_app_version(),
|
||||||
|
"provider_presets_json": json.dumps(PROVIDER_PRESETS),
|
||||||
|
"provider_models_json": "{}",
|
||||||
}
|
}
|
||||||
|
|
||||||
if not is_initializing:
|
if not is_initializing:
|
||||||
@@ -384,12 +394,14 @@ class ModelListingHandler:
|
|||||||
)
|
)
|
||||||
|
|
||||||
# View-local-versions filter: show all local versions of a specific model
|
# View-local-versions filter: show all local versions of a specific model
|
||||||
|
# Accepts either a CivitAI modelId (int) or a HF group key like "hf:user/repo"
|
||||||
civitai_model_id = request.query.get("civitai_model_id")
|
civitai_model_id = request.query.get("civitai_model_id")
|
||||||
if civitai_model_id is not None:
|
if civitai_model_id is not None:
|
||||||
try:
|
try:
|
||||||
civitai_model_id = int(civitai_model_id)
|
civitai_model_id = int(civitai_model_id)
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
civitai_model_id = None
|
# Keep as string — could be an HF group key (e.g. "hf:user/repo")
|
||||||
|
pass
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"page": page,
|
"page": page,
|
||||||
@@ -527,6 +539,7 @@ class ModelManagementHandler:
|
|||||||
# Update model_data with new hash
|
# Update model_data with new hash
|
||||||
model_data["sha256"] = sha256
|
model_data["sha256"] = sha256
|
||||||
model_data["hash_status"] = "completed"
|
model_data["hash_status"] = "completed"
|
||||||
|
hash_status = "completed"
|
||||||
else:
|
else:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": False, "error": "No SHA256 hash found"}, status=400
|
{"success": False, "error": "No SHA256 hash found"}, status=400
|
||||||
@@ -534,6 +547,32 @@ class ModelManagementHandler:
|
|||||||
|
|
||||||
await MetadataManager.hydrate_model_data(model_data)
|
await MetadataManager.hydrate_model_data(model_data)
|
||||||
|
|
||||||
|
# hydrate_model_data replaces model_data with .metadata.json content,
|
||||||
|
# which may lack sha256. Restore from cache and persist the fix.
|
||||||
|
if not model_data.get("sha256"):
|
||||||
|
if sha256:
|
||||||
|
model_data["sha256"] = sha256
|
||||||
|
model_data["hash_status"] = model_data.get("hash_status", hash_status)
|
||||||
|
data_to_save = model_data.copy()
|
||||||
|
data_to_save.pop("folder", None)
|
||||||
|
await MetadataManager.save_metadata(file_path, data_to_save)
|
||||||
|
else:
|
||||||
|
sha256 = await calculate_sha256(file_path)
|
||||||
|
if sha256:
|
||||||
|
model_data["sha256"] = sha256.lower()
|
||||||
|
model_data["hash_status"] = "completed"
|
||||||
|
data_to_save = model_data.copy()
|
||||||
|
data_to_save.pop("folder", None)
|
||||||
|
await MetadataManager.save_metadata(file_path, data_to_save)
|
||||||
|
else:
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": False,
|
||||||
|
"error": "Failed to compute SHA256 hash for model",
|
||||||
|
},
|
||||||
|
status=500,
|
||||||
|
)
|
||||||
|
|
||||||
success, error = await self._metadata_sync.fetch_and_update_model(
|
success, error = await self._metadata_sync.fetch_and_update_model(
|
||||||
sha256=model_data["sha256"],
|
sha256=model_data["sha256"],
|
||||||
file_path=file_path,
|
file_path=file_path,
|
||||||
@@ -556,7 +595,12 @@ class ModelManagementHandler:
|
|||||||
{"success": False, "error": OFFLINE_FRIENDLY_MESSAGE},
|
{"success": False, "error": OFFLINE_FRIENDLY_MESSAGE},
|
||||||
status=503,
|
status=503,
|
||||||
)
|
)
|
||||||
self._logger.error("Error fetching from CivitAI: %s", exc, exc_info=True)
|
self._logger.error(
|
||||||
|
"Error fetching from CivitAI for %s: %s",
|
||||||
|
locals().get("file_path", "unknown"),
|
||||||
|
exc,
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
async def relink_civitai(self, request: web.Request) -> web.Response:
|
async def relink_civitai(self, request: web.Request) -> web.Response:
|
||||||
@@ -963,6 +1007,8 @@ class ModelQueryHandler:
|
|||||||
limit = int(request.query.get("limit", "20"))
|
limit = int(request.query.get("limit", "20"))
|
||||||
if limit < 0:
|
if limit < 0:
|
||||||
limit = 20
|
limit = 20
|
||||||
|
elif limit > 200:
|
||||||
|
limit = 20
|
||||||
top_tags = await self._service.get_top_tags(limit)
|
top_tags = await self._service.get_top_tags(limit)
|
||||||
return web.json_response({"success": True, "tags": top_tags})
|
return web.json_response({"success": True, "tags": top_tags})
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
@@ -971,6 +1017,22 @@ class ModelQueryHandler:
|
|||||||
{"success": False, "error": "Internal server error"}, status=500
|
{"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:
|
async def get_base_models(self, request: web.Request) -> web.Response:
|
||||||
try:
|
try:
|
||||||
limit = int(request.query.get("limit", "20"))
|
limit = int(request.query.get("limit", "20"))
|
||||||
@@ -1265,9 +1327,13 @@ class ModelQueryHandler:
|
|||||||
text=f"{self._service.model_type.capitalize()} file name is required",
|
text=f"{self._service.model_type.capitalize()} file name is required",
|
||||||
status=400,
|
status=400,
|
||||||
)
|
)
|
||||||
notes = await self._service.get_model_notes(model_name)
|
result = await self._service.get_model_notes(model_name)
|
||||||
if notes is not None:
|
if result is not None:
|
||||||
return web.json_response({"success": True, "notes": notes})
|
return web.json_response({
|
||||||
|
"success": True,
|
||||||
|
"notes": result["notes"],
|
||||||
|
"file_path": result["file_path"],
|
||||||
|
})
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{
|
{
|
||||||
"success": False,
|
"success": False,
|
||||||
@@ -1303,9 +1369,20 @@ class ModelQueryHandler:
|
|||||||
}
|
}
|
||||||
if include_license_flags:
|
if include_license_flags:
|
||||||
model_data = await self._service.get_model_info_by_name(model_name)
|
model_data = await self._service.get_model_info_by_name(model_name)
|
||||||
license_flags = (model_data or {}).get("license_flags")
|
# Only return license_flags when real CivitAI model license
|
||||||
if license_flags is not None:
|
# data exists. This mirrors ModelModal's guard
|
||||||
response_payload["license_flags"] = int(license_flags)
|
# (modelData?.civitai?.model) so the preview tooltip never
|
||||||
|
# shows misleading license icons for HF or other models
|
||||||
|
# without actual license metadata.
|
||||||
|
civitai_data = (model_data or {}).get("civitai") or {}
|
||||||
|
has_license_data = (
|
||||||
|
isinstance(civitai_data, dict)
|
||||||
|
and isinstance(civitai_data.get("model"), dict)
|
||||||
|
)
|
||||||
|
if has_license_data:
|
||||||
|
license_flags = (model_data or {}).get("license_flags")
|
||||||
|
if license_flags is not None:
|
||||||
|
response_payload["license_flags"] = int(license_flags)
|
||||||
# Include the user's license icon style preference so the
|
# Include the user's license icon style preference so the
|
||||||
# ComfyUI tooltip can pick the right set without a separate
|
# ComfyUI tooltip can pick the right set without a separate
|
||||||
# API call.
|
# API call.
|
||||||
@@ -1762,14 +1839,20 @@ class ModelDownloadHandler:
|
|||||||
|
|
||||||
async def delete_download_history_item(self, request: web.Request) -> web.Response:
|
async def delete_download_history_item(self, request: web.Request) -> web.Response:
|
||||||
try:
|
try:
|
||||||
item_id = int(request.query.get("id", "0"))
|
download_id = request.query.get("download_id")
|
||||||
if not item_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(
|
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()
|
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})
|
return web.json_response({"success": deleted})
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
self._logger.error(
|
self._logger.error(
|
||||||
@@ -1779,14 +1862,20 @@ class ModelDownloadHandler:
|
|||||||
|
|
||||||
async def retry_download_from_history(self, request: web.Request) -> web.Response:
|
async def retry_download_from_history(self, request: web.Request) -> web.Response:
|
||||||
try:
|
try:
|
||||||
item_id = int(request.query.get("id", "0"))
|
download_id = request.query.get("download_id")
|
||||||
if not item_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(
|
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()
|
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:
|
if item is None:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": False, "error": "History item not found or not retryable"},
|
{"success": False, "error": "History item not found or not retryable"},
|
||||||
@@ -2910,6 +2999,7 @@ class ModelHandlerSet:
|
|||||||
"bulk_delete_models": self.management.bulk_delete_models,
|
"bulk_delete_models": self.management.bulk_delete_models,
|
||||||
"verify_duplicates": self.management.verify_duplicates,
|
"verify_duplicates": self.management.verify_duplicates,
|
||||||
"get_top_tags": self.query.get_top_tags,
|
"get_top_tags": self.query.get_top_tags,
|
||||||
|
"search_tags": self.query.search_tags,
|
||||||
"get_base_models": self.query.get_base_models,
|
"get_base_models": self.query.get_base_models,
|
||||||
"get_model_types": self.query.get_model_types,
|
"get_model_types": self.query.get_model_types,
|
||||||
"scan_models": self.query.scan_models,
|
"scan_models": self.query.scan_models,
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import mimetypes
|
import mimetypes
|
||||||
import urllib.parse
|
import urllib.parse
|
||||||
@@ -53,6 +54,7 @@ class PreviewHandler:
|
|||||||
|
|
||||||
if not resolved.is_file():
|
if not resolved.is_file():
|
||||||
logger.debug("Preview file not found at %s", str(resolved))
|
logger.debug("Preview file not found at %s", str(resolved))
|
||||||
|
asyncio.create_task(self._cleanup_stale_preview_url(normalized))
|
||||||
raise web.HTTPNotFound(text="Preview file not found")
|
raise web.HTTPNotFound(text="Preview file not found")
|
||||||
|
|
||||||
# aiohttp's FileResponse handles range requests, content headers, and
|
# aiohttp's FileResponse handles range requests, content headers, and
|
||||||
@@ -69,6 +71,35 @@ class PreviewHandler:
|
|||||||
resp.headers["Cache-Control"] = "public, max-age=86400"
|
resp.headers["Cache-Control"] = "public, max-age=86400"
|
||||||
return resp
|
return resp
|
||||||
|
|
||||||
|
async def _cleanup_stale_preview_url(self, normalized_preview_path: str) -> None:
|
||||||
|
"""Fire-and-forget: clear stale preview_url from all model caches.
|
||||||
|
|
||||||
|
When a preview file is no longer on disk, remove its reference from
|
||||||
|
every cached entry so subsequent list API responses return an empty
|
||||||
|
``preview_url``, letting the frontend show the no-preview placeholder.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
from ...services.service_registry import ServiceRegistry
|
||||||
|
|
||||||
|
for service_name in ("lora_scanner", "checkpoint_scanner", "embedding_scanner"):
|
||||||
|
scanner = ServiceRegistry.get_service_sync(service_name)
|
||||||
|
if scanner is None or not hasattr(scanner, "_cache"):
|
||||||
|
continue
|
||||||
|
cache = getattr(scanner, "_cache", None)
|
||||||
|
if cache is None or not hasattr(cache, "clear_preview_by_path"):
|
||||||
|
continue
|
||||||
|
cleared = await cache.clear_preview_by_path(normalized_preview_path)
|
||||||
|
if cleared and hasattr(scanner, "_persist_current_cache"):
|
||||||
|
await scanner._persist_current_cache()
|
||||||
|
logger.info(
|
||||||
|
"Cleared stale preview_url for %d %s entries (%s)",
|
||||||
|
cleared,
|
||||||
|
service_name,
|
||||||
|
normalized_preview_path,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.debug("Failed to clean up stale preview_url: %s", exc)
|
||||||
|
|
||||||
async def _stream_file(
|
async def _stream_file(
|
||||||
self, request: web.Request, path: Path
|
self, request: web.Request, path: Path
|
||||||
) -> web.StreamResponse:
|
) -> web.StreamResponse:
|
||||||
|
|||||||
@@ -72,6 +72,7 @@ class RecipeHandlerSet:
|
|||||||
"save_recipe": self.management.save_recipe,
|
"save_recipe": self.management.save_recipe,
|
||||||
"delete_recipe": self.management.delete_recipe,
|
"delete_recipe": self.management.delete_recipe,
|
||||||
"get_top_tags": self.query.get_top_tags,
|
"get_top_tags": self.query.get_top_tags,
|
||||||
|
"search_tags": self.query.search_tags,
|
||||||
"get_base_models": self.query.get_base_models,
|
"get_base_models": self.query.get_base_models,
|
||||||
"get_roots": self.query.get_roots,
|
"get_roots": self.query.get_roots,
|
||||||
"get_folders": self.query.get_folders,
|
"get_folders": self.query.get_folders,
|
||||||
@@ -317,12 +318,11 @@ class RecipeQueryHandler:
|
|||||||
raise RuntimeError("Recipe scanner unavailable")
|
raise RuntimeError("Recipe scanner unavailable")
|
||||||
|
|
||||||
limit = int(request.query.get("limit", "20"))
|
limit = int(request.query.get("limit", "20"))
|
||||||
cache = await recipe_scanner.get_cached_data()
|
if limit < 0:
|
||||||
|
limit = 20
|
||||||
tag_counts: Dict[str, int] = {}
|
elif limit > 200:
|
||||||
for recipe in getattr(cache, "raw_data", []):
|
limit = 20
|
||||||
for tag in recipe.get("tags", []) or []:
|
tag_counts = await self._get_recipe_tag_counts(recipe_scanner)
|
||||||
tag_counts[tag] = tag_counts.get(tag, 0) + 1
|
|
||||||
|
|
||||||
sorted_tags = [
|
sorted_tags = [
|
||||||
{"tag": tag, "count": count} for tag, count in tag_counts.items()
|
{"tag": tag, "count": count} for tag, count in tag_counts.items()
|
||||||
@@ -333,6 +333,55 @@ class RecipeQueryHandler:
|
|||||||
self._logger.error("Error retrieving top tags: %s", exc, exc_info=True)
|
self._logger.error("Error retrieving top tags: %s", exc, exc_info=True)
|
||||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
async def search_tags(self, request: web.Request) -> web.Response:
|
||||||
|
try:
|
||||||
|
await self._ensure_dependencies_ready()
|
||||||
|
recipe_scanner = self._recipe_scanner_getter()
|
||||||
|
if recipe_scanner is None:
|
||||||
|
raise RuntimeError("Recipe scanner unavailable")
|
||||||
|
|
||||||
|
query = request.query.get("q", "")
|
||||||
|
limit = int(request.query.get("limit", "20"))
|
||||||
|
if limit < 0:
|
||||||
|
limit = 20
|
||||||
|
elif limit > 200:
|
||||||
|
limit = 20
|
||||||
|
|
||||||
|
tag_counts = await self._get_recipe_tag_counts(recipe_scanner)
|
||||||
|
normalized_query = (query or "").strip().lower()
|
||||||
|
if not normalized_query:
|
||||||
|
sorted_tags = [
|
||||||
|
{"tag": tag, "count": count} for tag, count in tag_counts.items()
|
||||||
|
]
|
||||||
|
sorted_tags.sort(key=lambda entry: entry["count"], reverse=True)
|
||||||
|
return web.json_response(
|
||||||
|
{"success": True, "tags": sorted_tags[: (limit if limit > 0 else 20)]}
|
||||||
|
)
|
||||||
|
|
||||||
|
matched = [
|
||||||
|
{"tag": tag, "count": count}
|
||||||
|
for tag, count in tag_counts.items()
|
||||||
|
if normalized_query in tag.lower()
|
||||||
|
]
|
||||||
|
matched.sort(key=lambda entry: entry["count"], reverse=True)
|
||||||
|
if limit == 0:
|
||||||
|
result = matched
|
||||||
|
else:
|
||||||
|
result = matched[:limit]
|
||||||
|
return web.json_response({"success": True, "tags": result})
|
||||||
|
except Exception as exc:
|
||||||
|
self._logger.error("Error searching recipe tags: %s", exc, exc_info=True)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
async def _get_recipe_tag_counts(self, recipe_scanner) -> Dict[str, int]:
|
||||||
|
"""Compute tag->count mapping from cached recipe data."""
|
||||||
|
cache = await recipe_scanner.get_cached_data()
|
||||||
|
tag_counts: Dict[str, int] = {}
|
||||||
|
for recipe in getattr(cache, "raw_data", []):
|
||||||
|
for tag in recipe.get("tags", []) or []:
|
||||||
|
tag_counts[tag] = tag_counts.get(tag, 0) + 1
|
||||||
|
return tag_counts
|
||||||
|
|
||||||
async def get_base_models(self, request: web.Request) -> web.Response:
|
async def get_base_models(self, request: web.Request) -> web.Response:
|
||||||
try:
|
try:
|
||||||
await self._ensure_dependencies_ready()
|
await self._ensure_dependencies_ready()
|
||||||
|
|||||||
@@ -22,6 +22,8 @@ class RouteDefinition:
|
|||||||
MISC_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
MISC_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
||||||
RouteDefinition("GET", "/api/lm/settings", "get_settings"),
|
RouteDefinition("GET", "/api/lm/settings", "get_settings"),
|
||||||
RouteDefinition("POST", "/api/lm/settings", "update_settings"),
|
RouteDefinition("POST", "/api/lm/settings", "update_settings"),
|
||||||
|
RouteDefinition("GET", "/api/lm/llm/models", "get_llm_models"),
|
||||||
|
RouteDefinition("GET", "/api/lm/llm/provider-models", "get_provider_models"),
|
||||||
RouteDefinition("GET", "/api/lm/doctor/diagnostics", "get_doctor_diagnostics"),
|
RouteDefinition("GET", "/api/lm/doctor/diagnostics", "get_doctor_diagnostics"),
|
||||||
RouteDefinition("POST", "/api/lm/doctor/repair-cache", "repair_doctor_cache"),
|
RouteDefinition("POST", "/api/lm/doctor/repair-cache", "repair_doctor_cache"),
|
||||||
RouteDefinition("POST", "/api/lm/doctor/resolve-filename-conflicts", "resolve_doctor_filename_conflicts"),
|
RouteDefinition("POST", "/api/lm/doctor/resolve-filename-conflicts", "resolve_doctor_filename_conflicts"),
|
||||||
@@ -37,10 +39,12 @@ MISC_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
|||||||
RouteDefinition("POST", "/api/lm/update-usage-stats", "update_usage_stats"),
|
RouteDefinition("POST", "/api/lm/update-usage-stats", "update_usage_stats"),
|
||||||
RouteDefinition("GET", "/api/lm/get-usage-stats", "get_usage_stats"),
|
RouteDefinition("GET", "/api/lm/get-usage-stats", "get_usage_stats"),
|
||||||
RouteDefinition("POST", "/api/lm/update-lora-code", "update_lora_code"),
|
RouteDefinition("POST", "/api/lm/update-lora-code", "update_lora_code"),
|
||||||
|
RouteDefinition("GET", "/api/lm/update-lora-code", "get_update_lora_code"),
|
||||||
RouteDefinition("GET", "/api/lm/trained-words", "get_trained_words"),
|
RouteDefinition("GET", "/api/lm/trained-words", "get_trained_words"),
|
||||||
RouteDefinition("GET", "/api/lm/model-example-files", "get_model_example_files"),
|
RouteDefinition("GET", "/api/lm/model-example-files", "get_model_example_files"),
|
||||||
RouteDefinition("POST", "/api/lm/register-nodes", "register_nodes"),
|
RouteDefinition("POST", "/api/lm/register-nodes", "register_nodes"),
|
||||||
RouteDefinition("POST", "/api/lm/update-node-widget", "update_node_widget"),
|
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/get-registry", "get_registry"),
|
||||||
RouteDefinition("GET", "/api/lm/check-model-exists", "check_model_exists"),
|
RouteDefinition("GET", "/api/lm/check-model-exists", "check_model_exists"),
|
||||||
RouteDefinition("GET", "/api/lm/check-models-exist", "check_models_exist"),
|
RouteDefinition("GET", "/api/lm/check-models-exist", "check_models_exist"),
|
||||||
@@ -101,6 +105,19 @@ MISC_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
|||||||
RouteDefinition(
|
RouteDefinition(
|
||||||
"POST", "/api/lm/download-hf-model", "download_hf_model"
|
"POST", "/api/lm/download-hf-model", "download_hf_model"
|
||||||
),
|
),
|
||||||
|
RouteDefinition(
|
||||||
|
"POST", "/api/lm/set-hf-url", "set_hf_url"
|
||||||
|
),
|
||||||
|
# Agent skill endpoints
|
||||||
|
RouteDefinition(
|
||||||
|
"GET", "/api/lm/agent/skills", "get_agent_skills"
|
||||||
|
),
|
||||||
|
RouteDefinition(
|
||||||
|
"POST", "/api/lm/agent/execute/{skill_name}", "execute_agent_skill"
|
||||||
|
),
|
||||||
|
RouteDefinition(
|
||||||
|
"POST", "/api/lm/agent/cancel", "cancel_agent_skill"
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -40,6 +40,7 @@ from .handlers.misc_handlers import (
|
|||||||
)
|
)
|
||||||
from .handlers.base_model_handlers import BaseModelHandlerSet
|
from .handlers.base_model_handlers import BaseModelHandlerSet
|
||||||
from .handlers.hf_handlers import HfHandler
|
from .handlers.hf_handlers import HfHandler
|
||||||
|
from .handlers.agent_handlers import AgentHandler
|
||||||
from .misc_route_registrar import MiscRouteRegistrar
|
from .misc_route_registrar import MiscRouteRegistrar
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -138,6 +139,7 @@ class MiscRoutes:
|
|||||||
example_workflows = ExampleWorkflowsHandler()
|
example_workflows = ExampleWorkflowsHandler()
|
||||||
base_model = BaseModelHandlerSet()
|
base_model = BaseModelHandlerSet()
|
||||||
hf_handler = HfHandler()
|
hf_handler = HfHandler()
|
||||||
|
agent_handler = AgentHandler()
|
||||||
|
|
||||||
return self._handler_set_factory(
|
return self._handler_set_factory(
|
||||||
health=health,
|
health=health,
|
||||||
@@ -158,6 +160,7 @@ class MiscRoutes:
|
|||||||
example_workflows=example_workflows,
|
example_workflows=example_workflows,
|
||||||
base_model=base_model,
|
base_model=base_model,
|
||||||
hf_handler=hf_handler,
|
hf_handler=hf_handler,
|
||||||
|
agent_handler=agent_handler,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -46,6 +46,7 @@ COMMON_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
|||||||
"GET", "/api/lm/{prefix}/auto-organize-progress", "get_auto_organize_progress"
|
"GET", "/api/lm/{prefix}/auto-organize-progress", "get_auto_organize_progress"
|
||||||
),
|
),
|
||||||
RouteDefinition("GET", "/api/lm/{prefix}/top-tags", "get_top_tags"),
|
RouteDefinition("GET", "/api/lm/{prefix}/top-tags", "get_top_tags"),
|
||||||
|
RouteDefinition("GET", "/api/lm/{prefix}/search-tags", "search_tags"),
|
||||||
RouteDefinition("GET", "/api/lm/{prefix}/base-models", "get_base_models"),
|
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}/model-types", "get_model_types"),
|
||||||
RouteDefinition("GET", "/api/lm/{prefix}/scan", "scan_models"),
|
RouteDefinition("GET", "/api/lm/{prefix}/scan", "scan_models"),
|
||||||
|
|||||||
@@ -29,6 +29,7 @@ ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
|||||||
RouteDefinition("POST", "/api/lm/recipes/save", "save_recipe"),
|
RouteDefinition("POST", "/api/lm/recipes/save", "save_recipe"),
|
||||||
RouteDefinition("DELETE", "/api/lm/recipe/{recipe_id}", "delete_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/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/base-models", "get_base_models"),
|
||||||
RouteDefinition("GET", "/api/lm/recipes/roots", "get_roots"),
|
RouteDefinition("GET", "/api/lm/recipes/roots", "get_roots"),
|
||||||
RouteDefinition("GET", "/api/lm/recipes/folders", "get_folders"),
|
RouteDefinition("GET", "/api/lm/recipes/folders", "get_folders"),
|
||||||
|
|||||||
+305
-37
@@ -38,6 +38,84 @@ def _clean_excludes() -> List[str]:
|
|||||||
return excludes
|
return excludes
|
||||||
|
|
||||||
|
|
||||||
|
def _stage_preserved_items(plugin_root: str) -> tuple[str, list[str]]:
|
||||||
|
"""Move preserved user-data items to a temp directory outside *plugin_root*.
|
||||||
|
|
||||||
|
This ensures that ``git reset --hard``, ``git clean -fd``, and ZIP-based
|
||||||
|
replacement cannot touch these files even when ``-e`` exclusion patterns
|
||||||
|
are mishandled (e.g. on Windows where forward-slash patterns may not
|
||||||
|
match backslash-prefixed paths in some Git builds, or where file locks
|
||||||
|
prevent deletion/recreation).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``(backup_root, staged_names)``: the temp directory path and the
|
||||||
|
list of item names that were successfully moved.
|
||||||
|
"""
|
||||||
|
backup_root = tempfile.mkdtemp(prefix='lora_manager_update_')
|
||||||
|
staged: list[str] = []
|
||||||
|
for name in _PRESERVE_DIRS:
|
||||||
|
src = os.path.join(plugin_root, name)
|
||||||
|
if not os.path.lexists(src):
|
||||||
|
continue
|
||||||
|
dst = os.path.join(backup_root, name)
|
||||||
|
try:
|
||||||
|
shutil.move(src, dst)
|
||||||
|
staged.append(name)
|
||||||
|
logger.debug("Staged '%s' for update safety", name)
|
||||||
|
except OSError:
|
||||||
|
# ``shutil.move`` may fail on Windows if a file handle inside
|
||||||
|
# the directory is still open (e.g. a SQLite WAL file). Fall
|
||||||
|
# back to copy-then-remove.
|
||||||
|
logger.debug("Move failed for '%s', falling back to copy", name)
|
||||||
|
try:
|
||||||
|
if os.path.isdir(src) and not os.path.islink(src):
|
||||||
|
shutil.copytree(src, dst, symlinks=True)
|
||||||
|
shutil.rmtree(src, ignore_errors=True)
|
||||||
|
else:
|
||||||
|
shutil.copy2(src, dst)
|
||||||
|
os.remove(src)
|
||||||
|
staged.append(name)
|
||||||
|
logger.info("Copied (then removed) '%s' for update safety", name)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Could not stage '%s': %s (will rely on git -e / skip lists)", name, exc
|
||||||
|
)
|
||||||
|
return backup_root, staged
|
||||||
|
|
||||||
|
|
||||||
|
def _restore_preserved_items(plugin_root: str, backup_root: str, staged: list[str]) -> None:
|
||||||
|
"""Move staged items back from *backup_root* into *plugin_root*.
|
||||||
|
|
||||||
|
Any leftover placeholder at the destination (created by git checkout or
|
||||||
|
ZIP extraction) is removed before the move.
|
||||||
|
"""
|
||||||
|
for name in staged:
|
||||||
|
src = os.path.join(backup_root, name)
|
||||||
|
dst = os.path.join(plugin_root, name)
|
||||||
|
try:
|
||||||
|
if os.path.lexists(dst):
|
||||||
|
if os.path.isdir(dst) and not os.path.islink(dst):
|
||||||
|
shutil.rmtree(dst, ignore_errors=True)
|
||||||
|
else:
|
||||||
|
os.remove(dst)
|
||||||
|
shutil.move(src, dst)
|
||||||
|
logger.debug("Restored '%s' after update", name)
|
||||||
|
except OSError:
|
||||||
|
logger.debug("Move failed restoring '%s', falling back to copy", name)
|
||||||
|
try:
|
||||||
|
if os.path.isdir(src) and not os.path.islink(src):
|
||||||
|
shutil.copytree(src, dst, symlinks=True, dirs_exist_ok=True)
|
||||||
|
shutil.rmtree(src, ignore_errors=True)
|
||||||
|
else:
|
||||||
|
shutil.copy2(src, dst)
|
||||||
|
os.remove(src)
|
||||||
|
logger.info("Copied '%s' back after update", name)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("Failed to restore '%s': %s", name, exc)
|
||||||
|
shutil.rmtree(backup_root, ignore_errors=True)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class UpdateRoutes:
|
class UpdateRoutes:
|
||||||
"""Routes for handling plugin update checks"""
|
"""Routes for handling plugin update checks"""
|
||||||
|
|
||||||
@@ -47,6 +125,7 @@ class UpdateRoutes:
|
|||||||
app.router.add_get('/api/lm/check-updates', UpdateRoutes.check_updates)
|
app.router.add_get('/api/lm/check-updates', UpdateRoutes.check_updates)
|
||||||
app.router.add_get('/api/lm/version-info', UpdateRoutes.get_version_info)
|
app.router.add_get('/api/lm/version-info', UpdateRoutes.get_version_info)
|
||||||
app.router.add_post('/api/lm/perform-update', UpdateRoutes.perform_update)
|
app.router.add_post('/api/lm/perform-update', UpdateRoutes.perform_update)
|
||||||
|
app.router.add_post('/api/lm/switch-channel', UpdateRoutes.switch_channel)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def check_updates(request):
|
async def check_updates(request):
|
||||||
@@ -65,10 +144,17 @@ class UpdateRoutes:
|
|||||||
|
|
||||||
# Fetch remote version from GitHub
|
# Fetch remote version from GitHub
|
||||||
if nightly:
|
if nightly:
|
||||||
remote_version, changelog = await UpdateRoutes._get_nightly_version()
|
local_hash = git_info.get('short_hash', '')
|
||||||
releases = None
|
nightly_version, releases_result = await asyncio.gather(
|
||||||
|
UpdateRoutes._get_nightly_version(local_hash),
|
||||||
|
UpdateRoutes._get_remote_version()
|
||||||
|
)
|
||||||
|
remote_version, _, behind_by, commit_date = nightly_version
|
||||||
|
_, changelog, releases = releases_result
|
||||||
else:
|
else:
|
||||||
remote_version, changelog, releases = await UpdateRoutes._get_remote_version()
|
remote_version, changelog, releases = await UpdateRoutes._get_remote_version()
|
||||||
|
behind_by = 0
|
||||||
|
commit_date = ''
|
||||||
|
|
||||||
# Compare versions
|
# Compare versions
|
||||||
if nightly:
|
if nightly:
|
||||||
@@ -81,6 +167,10 @@ class UpdateRoutes:
|
|||||||
remote_version.replace('v', '')
|
remote_version.replace('v', '')
|
||||||
)
|
)
|
||||||
|
|
||||||
|
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
plugin_root = os.path.dirname(os.path.dirname(current_dir))
|
||||||
|
has_git = os.path.exists(os.path.join(plugin_root, '.git'))
|
||||||
|
|
||||||
response_data = {
|
response_data = {
|
||||||
'success': True,
|
'success': True,
|
||||||
'current_version': local_version,
|
'current_version': local_version,
|
||||||
@@ -88,13 +178,13 @@ class UpdateRoutes:
|
|||||||
'update_available': update_available,
|
'update_available': update_available,
|
||||||
'changelog': changelog,
|
'changelog': changelog,
|
||||||
'git_info': git_info,
|
'git_info': git_info,
|
||||||
'nightly': nightly
|
'nightly': nightly,
|
||||||
|
'has_git': has_git,
|
||||||
|
'releases': releases,
|
||||||
|
'behind_by': behind_by,
|
||||||
|
'commit_date': commit_date
|
||||||
}
|
}
|
||||||
|
|
||||||
# Include releases list for stable mode
|
|
||||||
if releases is not None:
|
|
||||||
response_data['releases'] = releases
|
|
||||||
|
|
||||||
return web.json_response(response_data)
|
return web.json_response(response_data)
|
||||||
|
|
||||||
except NETWORK_EXCEPTIONS as e:
|
except NETWORK_EXCEPTIONS as e:
|
||||||
@@ -126,9 +216,14 @@ class UpdateRoutes:
|
|||||||
# Format: version-short_hash
|
# Format: version-short_hash
|
||||||
version_string = f"{local_version}-{short_hash}"
|
version_string = f"{local_version}-{short_hash}"
|
||||||
|
|
||||||
|
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
plugin_root = os.path.dirname(os.path.dirname(current_dir))
|
||||||
|
has_git = os.path.exists(os.path.join(plugin_root, '.git'))
|
||||||
|
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
'success': True,
|
'success': True,
|
||||||
'version': version_string
|
'version': version_string,
|
||||||
|
'has_git': has_git
|
||||||
})
|
})
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -156,20 +251,22 @@ class UpdateRoutes:
|
|||||||
if os.path.exists(settings_path):
|
if os.path.exists(settings_path):
|
||||||
with open(settings_path, 'r', encoding='utf-8') as f:
|
with open(settings_path, 'r', encoding='utf-8') as f:
|
||||||
settings_backup = f.read()
|
settings_backup = f.read()
|
||||||
logger.info("Backed up settings.json")
|
logger.debug("Backed up settings.json (%d bytes)", len(settings_backup))
|
||||||
|
|
||||||
git_folder = os.path.join(plugin_root, '.git')
|
staged_backup_dir, staged_items = _stage_preserved_items(plugin_root)
|
||||||
if os.path.exists(git_folder):
|
try:
|
||||||
# Git update
|
git_folder = os.path.join(plugin_root, '.git')
|
||||||
success, new_version = await UpdateRoutes._perform_git_update(plugin_root, nightly)
|
if os.path.exists(git_folder):
|
||||||
else:
|
success, new_version = await UpdateRoutes._perform_git_update(plugin_root, nightly)
|
||||||
# Fallback: Download ZIP and replace files
|
else:
|
||||||
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
||||||
|
finally:
|
||||||
|
_restore_preserved_items(plugin_root, staged_backup_dir, staged_items)
|
||||||
|
|
||||||
if settings_backup and success:
|
if settings_backup and success:
|
||||||
with open(settings_path, 'w', encoding='utf-8') as f:
|
with open(settings_path, 'w', encoding='utf-8') as f:
|
||||||
f.write(settings_backup)
|
f.write(settings_backup)
|
||||||
logger.info("Restored settings.json")
|
logger.debug("Restored settings.json content (%d bytes)", len(settings_backup))
|
||||||
|
|
||||||
if success:
|
if success:
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
@@ -190,6 +287,164 @@ class UpdateRoutes:
|
|||||||
'error': str(e)
|
'error': str(e)
|
||||||
})
|
})
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def switch_channel(request):
|
||||||
|
"""
|
||||||
|
Switch between release and nightly update channels.
|
||||||
|
|
||||||
|
ZIP/CNR install → Nightly: git init + checkout main (one-way upgrade)
|
||||||
|
Git install → Release: git checkout latest tag (.git preserved)
|
||||||
|
ZIP/CNR install → Release: ZIP download (no .git, stays in ZIP mode)
|
||||||
|
Git install → Nightly: git checkout main + pull
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
body = await request.json() if request.has_body else {}
|
||||||
|
channel = body.get('channel', '')
|
||||||
|
|
||||||
|
if channel not in ('release', 'nightly'):
|
||||||
|
return web.json_response({
|
||||||
|
'success': False,
|
||||||
|
'error': f'Invalid channel: {channel}. Must be "release" or "nightly".'
|
||||||
|
})
|
||||||
|
|
||||||
|
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
plugin_root = os.path.dirname(os.path.dirname(current_dir))
|
||||||
|
|
||||||
|
settings_path = ensure_settings_file(logger)
|
||||||
|
settings_backup = None
|
||||||
|
if os.path.exists(settings_path):
|
||||||
|
with open(settings_path, 'r', encoding='utf-8') as f:
|
||||||
|
settings_backup = f.read()
|
||||||
|
logger.debug("Backed up settings.json before channel switch (%d bytes)", len(settings_backup))
|
||||||
|
|
||||||
|
staged_backup_dir, staged_items = _stage_preserved_items(plugin_root)
|
||||||
|
try:
|
||||||
|
git_folder = os.path.join(plugin_root, '.git')
|
||||||
|
|
||||||
|
if channel == 'nightly':
|
||||||
|
git_backup = None
|
||||||
|
if os.path.exists(git_folder):
|
||||||
|
git_backup = UpdateRoutes._backup_git(git_folder, 'nightly')
|
||||||
|
|
||||||
|
success = False
|
||||||
|
new_version = ''
|
||||||
|
try:
|
||||||
|
if os.path.exists(git_folder):
|
||||||
|
success, new_version = await UpdateRoutes._perform_git_update(
|
||||||
|
plugin_root, nightly=True
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
success, new_version = UpdateRoutes._init_git_repo(plugin_root)
|
||||||
|
finally:
|
||||||
|
UpdateRoutes._restore_git(git_backup, git_folder, success, 'nightly')
|
||||||
|
else:
|
||||||
|
success = False
|
||||||
|
new_version = ''
|
||||||
|
if os.path.exists(git_folder):
|
||||||
|
success, new_version = await UpdateRoutes._perform_git_update(
|
||||||
|
plugin_root, nightly=False
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
tracking_file = os.path.join(plugin_root, '.tracking')
|
||||||
|
if os.path.exists(tracking_file):
|
||||||
|
os.remove(tracking_file)
|
||||||
|
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
||||||
|
finally:
|
||||||
|
_restore_preserved_items(plugin_root, staged_backup_dir, staged_items)
|
||||||
|
|
||||||
|
if settings_backup and success:
|
||||||
|
with open(settings_path, 'w', encoding='utf-8') as f:
|
||||||
|
f.write(settings_backup)
|
||||||
|
logger.debug("Restored settings.json content after channel switch (%d bytes)", len(settings_backup))
|
||||||
|
|
||||||
|
if success:
|
||||||
|
return web.json_response({
|
||||||
|
'success': True,
|
||||||
|
'channel': channel,
|
||||||
|
'new_version': new_version,
|
||||||
|
'message': f'Switched to {channel} channel'
|
||||||
|
})
|
||||||
|
else:
|
||||||
|
return web.json_response({
|
||||||
|
'success': False,
|
||||||
|
'error': f'Failed to switch to {channel} channel'
|
||||||
|
})
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Failed to switch channel: %s", e, exc_info=True)
|
||||||
|
return web.json_response({
|
||||||
|
'success': False,
|
||||||
|
'error': str(e)
|
||||||
|
})
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _init_git_repo(plugin_root: str) -> tuple[bool, str]:
|
||||||
|
"""
|
||||||
|
Initialize a Git repository in a ZIP-installed plugin folder.
|
||||||
|
Clones the remote history and checks out main branch.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
import git
|
||||||
|
except ImportError:
|
||||||
|
logger.error(
|
||||||
|
"GitPython is not available: cannot initialize git repo. "
|
||||||
|
"Install git or set $GIT_PYTHON_GIT_EXECUTABLE to the git binary path."
|
||||||
|
)
|
||||||
|
return False, ""
|
||||||
|
|
||||||
|
clean_excludes = _clean_excludes()
|
||||||
|
|
||||||
|
try:
|
||||||
|
repo = git.Repo.init(plugin_root)
|
||||||
|
origin = repo.create_remote(
|
||||||
|
'origin',
|
||||||
|
'https://github.com/willmiao/ComfyUI-Lora-Manager.git'
|
||||||
|
)
|
||||||
|
origin.fetch()
|
||||||
|
|
||||||
|
repo.create_head('main', origin.refs.main)
|
||||||
|
repo.git.checkout('main', '--force')
|
||||||
|
repo.git.reset('--hard')
|
||||||
|
repo.git.clean('-fd', *clean_excludes)
|
||||||
|
|
||||||
|
tracking_file = os.path.join(plugin_root, '.tracking')
|
||||||
|
if os.path.exists(tracking_file):
|
||||||
|
os.remove(tracking_file)
|
||||||
|
logger.info("Removed .tracking file (now in git mode)")
|
||||||
|
|
||||||
|
new_version = f"main-{repo.head.commit.hexsha[:7]}"
|
||||||
|
logger.info("Initialized git repo on main branch: %s", new_version)
|
||||||
|
return True, new_version
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Failed to initialize git repo: %s", e, exc_info=True)
|
||||||
|
return False, ""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _backup_git(git_folder, label):
|
||||||
|
try:
|
||||||
|
backup_dir = tempfile.mkdtemp()
|
||||||
|
backup = os.path.join(backup_dir, '.git')
|
||||||
|
shutil.copytree(git_folder, backup)
|
||||||
|
logger.info("Backed up .git before switching to %s", label)
|
||||||
|
return backup
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Failed to backup .git before %s switch: %s", label, e)
|
||||||
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _restore_git(git_backup, git_folder, success, label):
|
||||||
|
if git_backup and not success:
|
||||||
|
try:
|
||||||
|
if os.path.exists(git_folder):
|
||||||
|
shutil.rmtree(git_folder)
|
||||||
|
shutil.copytree(git_backup, git_folder)
|
||||||
|
logger.info("Restored .git after failed %s switch", label)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Failed to restore .git after %s switch: %s", label, e)
|
||||||
|
if git_backup:
|
||||||
|
shutil.rmtree(os.path.dirname(git_backup), ignore_errors=True)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def _download_and_replace_zip(plugin_root: str) -> tuple[bool, str]:
|
async def _download_and_replace_zip(plugin_root: str) -> tuple[bool, str]:
|
||||||
"""
|
"""
|
||||||
@@ -244,8 +499,7 @@ class UpdateRoutes:
|
|||||||
except Exception:
|
except Exception:
|
||||||
logger.debug("Could not close downloaded-version history database", exc_info=True)
|
logger.debug("Could not close downloaded-version history database", exc_info=True)
|
||||||
|
|
||||||
# Skip settings.json, civitai, model cache and runtime cache folders
|
UpdateRoutes._clean_plugin_folder(plugin_root, skip_files=list(_PRESERVE_DIRS))
|
||||||
UpdateRoutes._clean_plugin_folder(plugin_root, skip_files=['settings.json', 'civitai', 'model_cache', 'cache', 'wildcards', 'backups', 'stats'])
|
|
||||||
|
|
||||||
# Extract ZIP to temp dir
|
# Extract ZIP to temp dir
|
||||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||||
@@ -255,7 +509,7 @@ class UpdateRoutes:
|
|||||||
extracted_root = next(os.scandir(tmp_dir)).path
|
extracted_root = next(os.scandir(tmp_dir)).path
|
||||||
|
|
||||||
# Copy files, skipping user data that should be preserved
|
# Copy files, skipping user data that should be preserved
|
||||||
skip_items = {'settings.json', 'civitai', 'wildcards', 'backups', 'stats'}
|
skip_items = set(_PRESERVE_DIRS)
|
||||||
for item in os.listdir(extracted_root):
|
for item in os.listdir(extracted_root):
|
||||||
if item in skip_items:
|
if item in skip_items:
|
||||||
continue
|
continue
|
||||||
@@ -272,7 +526,7 @@ class UpdateRoutes:
|
|||||||
# for ComfyUI Manager to work properly
|
# for ComfyUI Manager to work properly
|
||||||
tracking_info_file = os.path.join(plugin_root, '.tracking')
|
tracking_info_file = os.path.join(plugin_root, '.tracking')
|
||||||
tracking_files = []
|
tracking_files = []
|
||||||
skip_tracked = {'civitai', 'wildcards', 'backups', 'stats'}
|
skip_tracked = set(_PRESERVE_DIRS) - {'settings.json'}
|
||||||
for root, dirs, files in os.walk(extracted_root):
|
for root, dirs, files in os.walk(extracted_root):
|
||||||
# Skip user data directories and their contents
|
# Skip user data directories and their contents
|
||||||
rel_root = os.path.relpath(root, extracted_root)
|
rel_root = os.path.relpath(root, extracted_root)
|
||||||
@@ -296,6 +550,7 @@ class UpdateRoutes:
|
|||||||
logger.error(f"ZIP update failed: {e}", exc_info=True)
|
logger.error(f"ZIP update failed: {e}", exc_info=True)
|
||||||
return False, ""
|
return False, ""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
def _clean_plugin_folder(plugin_root, skip_files=None):
|
def _clean_plugin_folder(plugin_root, skip_files=None):
|
||||||
skip_files = skip_files or []
|
skip_files = skip_files or []
|
||||||
for item in os.listdir(plugin_root):
|
for item in os.listdir(plugin_root):
|
||||||
@@ -308,41 +563,54 @@ class UpdateRoutes:
|
|||||||
os.remove(path)
|
os.remove(path)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def _get_nightly_version() -> tuple[str, List[str]]:
|
async def _get_nightly_version(local_hash: str = "") -> tuple[str, List[str], int, str]:
|
||||||
"""
|
|
||||||
Fetch latest commit from main branch
|
|
||||||
"""
|
|
||||||
repo_owner = "willmiao"
|
repo_owner = "willmiao"
|
||||||
repo_name = "ComfyUI-Lora-Manager"
|
repo_name = "ComfyUI-Lora-Manager"
|
||||||
|
|
||||||
# Use GitHub API to fetch the latest commit from main branch
|
|
||||||
github_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/commits/main"
|
github_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/commits/main"
|
||||||
|
|
||||||
try:
|
try:
|
||||||
downloader = await get_downloader()
|
downloader = await get_downloader()
|
||||||
success, data = await downloader.make_request('GET', github_url, custom_headers={'Accept': 'application/vnd.github+json'})
|
success, data = await downloader.make_request(
|
||||||
|
'GET', github_url,
|
||||||
|
custom_headers={'Accept': 'application/vnd.github+json'}
|
||||||
|
)
|
||||||
|
|
||||||
if not success:
|
if not success:
|
||||||
logger.warning(f"Failed to fetch GitHub commit: {data}")
|
logger.warning("Failed to fetch GitHub commit: %s", data)
|
||||||
return "main", []
|
return "main", [], 0, ""
|
||||||
|
|
||||||
commit_sha = data.get('sha', '')[:7] # Short hash
|
commit_sha = data.get('sha', '')[:7]
|
||||||
commit_message = data.get('commit', {}).get('message', '')
|
commit_message = data.get('commit', {}).get('message', '')
|
||||||
|
commit_date = data.get('commit', {}).get('committer', {}).get('date', '')[:10]
|
||||||
|
|
||||||
# Format as "main-{short_hash}"
|
|
||||||
version = f"main-{commit_sha}"
|
version = f"main-{commit_sha}"
|
||||||
|
|
||||||
# Use commit message as changelog
|
|
||||||
changelog = [commit_message] if commit_message else []
|
changelog = [commit_message] if commit_message else []
|
||||||
|
|
||||||
return version, changelog
|
behind_by = 0
|
||||||
|
if local_hash and local_hash not in ('unknown', 'stable'):
|
||||||
|
compare_url = (
|
||||||
|
f"https://api.github.com/repos/{repo_owner}/{repo_name}"
|
||||||
|
f"/compare/{local_hash}...main"
|
||||||
|
)
|
||||||
|
c_ok, c_data = await downloader.make_request(
|
||||||
|
'GET', compare_url,
|
||||||
|
custom_headers={'Accept': 'application/vnd.github+json'}
|
||||||
|
)
|
||||||
|
if c_ok:
|
||||||
|
if c_data.get('status') in ('ahead', 'diverged'):
|
||||||
|
behind_by = c_data.get('ahead_by', 0)
|
||||||
|
else:
|
||||||
|
behind_by = c_data.get('behind_by', 0)
|
||||||
|
|
||||||
|
return version, changelog, behind_by, commit_date
|
||||||
|
|
||||||
except NETWORK_EXCEPTIONS as e:
|
except NETWORK_EXCEPTIONS as e:
|
||||||
logger.warning("Unable to reach GitHub for nightly version: %s", e)
|
logger.warning("Unable to reach GitHub for nightly version: %s", e)
|
||||||
return "main", []
|
return "main", [], 0, ""
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error fetching nightly version: {e}", exc_info=True)
|
logger.error("Error fetching nightly version: %s", e, exc_info=True)
|
||||||
return "main", []
|
return "main", [], 0, ""
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _compare_nightly_versions(local_git_info: Dict[str, str], remote_version: str) -> bool:
|
def _compare_nightly_versions(local_git_info: Dict[str, str], remote_version: str) -> bool:
|
||||||
|
|||||||
@@ -0,0 +1,27 @@
|
|||||||
|
"""LLM-powered metadata enrichment pipeline infrastructure.
|
||||||
|
|
||||||
|
This package provides the orchestration layer for LLM-powered features.
|
||||||
|
Skills define *what* to do (prompt template). The :class:`AgentService`
|
||||||
|
handles *how* (LLM calls, context gathering, validation, progress).
|
||||||
|
|
||||||
|
NOTE: The current implementation is a code-driven pipeline, not a true
|
||||||
|
agent loop. Future agent orchestration (LLM-driven tool selection) will
|
||||||
|
live alongside this package with its own namespace.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from .skill_definition import SkillDefinition, SkillPermissions
|
||||||
|
from .skill_registry import SkillRegistry
|
||||||
|
from .agent_service import AgentService, AgentProgressReporter, SkillResult
|
||||||
|
from .post_processor import PostProcessor
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"AgentProgressReporter",
|
||||||
|
"AgentService",
|
||||||
|
"PostProcessor",
|
||||||
|
"SkillDefinition",
|
||||||
|
"SkillPermissions",
|
||||||
|
"SkillRegistry",
|
||||||
|
"SkillResult",
|
||||||
|
]
|
||||||
@@ -0,0 +1,489 @@
|
|||||||
|
"""Pipeline orchestration service.
|
||||||
|
|
||||||
|
The :class:`AgentService` coordinates LLM-powered pipeline execution:
|
||||||
|
|
||||||
|
1. Look up the pipeline definition in :class:`SkillRegistry`
|
||||||
|
2. Validate input against its ``input_schema``
|
||||||
|
3. Prepare context via :mod:`~py.metadata_ops` (read metadata, list base models, fetch HF README)
|
||||||
|
4. If ``llm_required``: call :class:`LLMService` with the rendered prompt
|
||||||
|
5. Post-process via :class:`PostProcessor` (delegates I/O to :mod:`~py.metadata_ops`)
|
||||||
|
6. Broadcast progress and completion via :class:`WebSocketManager`
|
||||||
|
|
||||||
|
Pipeline definitions (*skills*) describe *what* to do (prompt template).
|
||||||
|
The AgentService handles *how* (LLM calls, context gathering, validation,
|
||||||
|
progress).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
import aiohttp
|
||||||
|
|
||||||
|
import os
|
||||||
|
|
||||||
|
from ...config import config
|
||||||
|
from ..llm_service import LLMService
|
||||||
|
from ..websocket_manager import ws_manager
|
||||||
|
from .post_processor import PostProcessor
|
||||||
|
from .skill_registry import SkillRegistry
|
||||||
|
from .skills.enrich_hf_metadata.readme_processor import (
|
||||||
|
clean_readme_for_llm,
|
||||||
|
extract_relevant_section,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class AgentProgressReporter:
|
||||||
|
"""Protocol-compatible progress reporter backed by WebSocket broadcast."""
|
||||||
|
|
||||||
|
async def on_progress(self, payload: Dict[str, Any]) -> None:
|
||||||
|
await ws_manager.broadcast(payload)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class SkillResult:
|
||||||
|
"""Outcome of a skill execution."""
|
||||||
|
|
||||||
|
success: bool
|
||||||
|
updated_models: List[Dict[str, Any]] = field(default_factory=list)
|
||||||
|
errors: List[str] = field(default_factory=list)
|
||||||
|
summary: str = ""
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_schema(data: Any, schema: Dict[str, Any], path: str = "") -> List[str]:
|
||||||
|
"""Minimal JSON schema validator.
|
||||||
|
|
||||||
|
Supports a subset of JSON Schema: ``type``, ``properties``, ``required``,
|
||||||
|
``items``, ``enum``. Returns a list of error messages (empty = valid).
|
||||||
|
"""
|
||||||
|
|
||||||
|
errors: List[str] = []
|
||||||
|
if not schema:
|
||||||
|
return errors
|
||||||
|
|
||||||
|
expected_type = schema.get("type")
|
||||||
|
if expected_type:
|
||||||
|
type_map = {
|
||||||
|
"string": str,
|
||||||
|
"number": (int, float),
|
||||||
|
"integer": int,
|
||||||
|
"boolean": bool,
|
||||||
|
"array": list,
|
||||||
|
"object": dict,
|
||||||
|
"null": type(None),
|
||||||
|
}
|
||||||
|
expected_py = type_map.get(expected_type)
|
||||||
|
if expected_py is not None and not isinstance(data, expected_py):
|
||||||
|
errors.append(f"{path or 'root'}: expected {expected_type}, got {type(data).__name__}")
|
||||||
|
return errors
|
||||||
|
|
||||||
|
if expected_type == "object" and isinstance(data, dict):
|
||||||
|
properties = schema.get("properties", {})
|
||||||
|
required = schema.get("required", [])
|
||||||
|
for req_key in required:
|
||||||
|
if req_key not in data:
|
||||||
|
errors.append(f"{path or 'root'}: missing required property '{req_key}'")
|
||||||
|
for key, value in data.items():
|
||||||
|
if key in properties:
|
||||||
|
errors.extend(_validate_schema(value, properties[key], f"{path}.{key}"))
|
||||||
|
|
||||||
|
if expected_type == "array" and isinstance(data, list):
|
||||||
|
items_schema = schema.get("items")
|
||||||
|
if items_schema:
|
||||||
|
for i, item in enumerate(data):
|
||||||
|
errors.extend(_validate_schema(item, items_schema, f"{path}[{i}]"))
|
||||||
|
|
||||||
|
if "enum" in schema and data not in schema["enum"]:
|
||||||
|
errors.append(f"{path or 'root'}: value '{data}' not in enum {schema['enum']}")
|
||||||
|
|
||||||
|
return errors
|
||||||
|
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Prompt template rendering
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _render_prompt(template: str, variables: Dict[str, Any]) -> str:
|
||||||
|
"""Render a prompt template with ``{{variable}}`` placeholders.
|
||||||
|
|
||||||
|
Uses simple regex substitution — no Jinja2 dependency needed.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def replace(match: re.Match) -> str:
|
||||||
|
key = match.group(1).strip()
|
||||||
|
value = variables.get(key, "")
|
||||||
|
if isinstance(value, (dict, list)):
|
||||||
|
return json.dumps(value, ensure_ascii=False, indent=2)
|
||||||
|
return str(value)
|
||||||
|
|
||||||
|
return re.sub(r"\{\{(\w+)\}\}", replace, template)
|
||||||
|
|
||||||
|
|
||||||
|
class AgentService:
|
||||||
|
"""Orchestrate agent skill execution.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
service = await AgentService.get_instance()
|
||||||
|
result = await service.execute_skill(
|
||||||
|
skill_name="enrich_hf_metadata",
|
||||||
|
input_data={"model_paths": ["/path/to/model.safetensors"]},
|
||||||
|
progress_callback=AgentProgressReporter(),
|
||||||
|
)
|
||||||
|
"""
|
||||||
|
|
||||||
|
_instance: Optional["AgentService"] = None
|
||||||
|
_lock: asyncio.Lock = asyncio.Lock()
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
skill_registry: Optional[SkillRegistry] = None,
|
||||||
|
llm_service: Optional[LLMService] = None,
|
||||||
|
) -> None:
|
||||||
|
self._registry = skill_registry
|
||||||
|
self._llm_service = llm_service
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def get_instance(cls) -> "AgentService":
|
||||||
|
"""Return the lazily-initialised global ``AgentService``."""
|
||||||
|
|
||||||
|
if cls._instance is None:
|
||||||
|
async with cls._lock:
|
||||||
|
if cls._instance is None:
|
||||||
|
cls._instance = cls(
|
||||||
|
skill_registry=await SkillRegistry.get_instance(),
|
||||||
|
llm_service=await LLMService.get_instance(),
|
||||||
|
)
|
||||||
|
return cls._instance
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def reset_instance(cls) -> None:
|
||||||
|
"""Reset the cached singleton — primarily for tests."""
|
||||||
|
|
||||||
|
cls._instance = None
|
||||||
|
|
||||||
|
async def _ensure_registry(self) -> SkillRegistry:
|
||||||
|
if self._registry is None:
|
||||||
|
self._registry = await SkillRegistry.get_instance()
|
||||||
|
return self._registry
|
||||||
|
|
||||||
|
async def _ensure_llm(self) -> LLMService:
|
||||||
|
if self._llm_service is None:
|
||||||
|
self._llm_service = await LLMService.get_instance()
|
||||||
|
return self._llm_service
|
||||||
|
|
||||||
|
async def list_skills(self) -> List[Dict[str, Any]]:
|
||||||
|
"""Return a JSON-serialisable list of available skills."""
|
||||||
|
|
||||||
|
registry = await self._ensure_registry()
|
||||||
|
return [
|
||||||
|
{
|
||||||
|
"name": s.name,
|
||||||
|
"title": s.title,
|
||||||
|
"description": s.description,
|
||||||
|
"llm_required": s.llm_required,
|
||||||
|
"model_type_filter": s.model_type_filter,
|
||||||
|
}
|
||||||
|
for s in registry.list_skills()
|
||||||
|
]
|
||||||
|
|
||||||
|
async def execute_skill(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
skill_name: str,
|
||||||
|
input_data: Dict[str, Any],
|
||||||
|
progress_callback: Optional[AgentProgressReporter] = None,
|
||||||
|
) -> SkillResult:
|
||||||
|
"""Execute a pipeline (skill) on the given models.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
skill_name: Name of the pipeline to execute
|
||||||
|
input_data: Input validated against the pipeline's ``input_schema``
|
||||||
|
progress_callback: Optional WebSocket progress reporter
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
:class:`SkillResult` with success status and updated model info
|
||||||
|
"""
|
||||||
|
|
||||||
|
registry = await self._ensure_registry()
|
||||||
|
skill = registry.get_skill(skill_name)
|
||||||
|
if skill is None:
|
||||||
|
return SkillResult(
|
||||||
|
success=False,
|
||||||
|
errors=[f"Skill not found: {skill_name}"],
|
||||||
|
summary=f"Skill '{skill_name}' does not exist",
|
||||||
|
)
|
||||||
|
|
||||||
|
input_errors = _validate_schema(input_data, skill.input_schema)
|
||||||
|
if input_errors:
|
||||||
|
return SkillResult(
|
||||||
|
success=False,
|
||||||
|
errors=input_errors,
|
||||||
|
summary=f"Invalid input: {'; '.join(input_errors)}",
|
||||||
|
)
|
||||||
|
|
||||||
|
model_paths = input_data.get("model_paths", [])
|
||||||
|
if not model_paths:
|
||||||
|
return SkillResult(
|
||||||
|
success=False,
|
||||||
|
errors=["No model_paths provided"],
|
||||||
|
summary="No models to process",
|
||||||
|
)
|
||||||
|
|
||||||
|
total = len(model_paths)
|
||||||
|
processed = 0
|
||||||
|
success_count = 0
|
||||||
|
skipped_count = 0
|
||||||
|
updated_models: List[Dict[str, Any]] = []
|
||||||
|
errors: List[str] = []
|
||||||
|
post_processor = PostProcessor()
|
||||||
|
|
||||||
|
await self._emit_progress(
|
||||||
|
progress_callback, skill_name, status="started",
|
||||||
|
total=total, processed=0, success=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
llm = await self._ensure_llm()
|
||||||
|
llm_configured = llm.is_configured() if skill.llm_required else True
|
||||||
|
|
||||||
|
for model_path in model_paths:
|
||||||
|
model_filename = os.path.basename(model_path)
|
||||||
|
logger.info(
|
||||||
|
"[%s] [%d/%d] %s",
|
||||||
|
skill_name, processed + 1, total, model_filename,
|
||||||
|
)
|
||||||
|
updated_data: Dict[str, Any] = {}
|
||||||
|
skip_model = False
|
||||||
|
try:
|
||||||
|
from ...metadata_ops import read_metadata
|
||||||
|
metadata = await read_metadata(model_path)
|
||||||
|
|
||||||
|
# Fast-fail: enrich_hf_metadata requires hf_url to have HF README context
|
||||||
|
if skill_name == "enrich_hf_metadata" and not metadata.get("hf_url", ""):
|
||||||
|
logger.info(
|
||||||
|
"[%s] SKIP %s — no hf_url in metadata",
|
||||||
|
skill_name, model_filename,
|
||||||
|
)
|
||||||
|
skipped_count += 1
|
||||||
|
skip_model = True
|
||||||
|
|
||||||
|
if not skip_model:
|
||||||
|
prompt_vars: Dict[str, Any] = {"model_path": model_path}
|
||||||
|
if skill.llm_required and llm_configured:
|
||||||
|
prompt_vars = await self._build_prompt_context(
|
||||||
|
skill_name, model_path, metadata, registry, llm,
|
||||||
|
)
|
||||||
|
|
||||||
|
llm_response: Optional[Dict[str, Any]] = None
|
||||||
|
if skill.llm_required and llm_configured:
|
||||||
|
prompt_template = registry.load_prompt(skill_name)
|
||||||
|
rendered = _render_prompt(prompt_template, prompt_vars)
|
||||||
|
llm_response = await llm.chat_completion_json(
|
||||||
|
system_prompt=prompt_vars.get(
|
||||||
|
"system_prompt",
|
||||||
|
"You are a helpful assistant that extracts structured metadata.",
|
||||||
|
),
|
||||||
|
user_prompt=rendered,
|
||||||
|
)
|
||||||
|
if llm_response:
|
||||||
|
logger.info(
|
||||||
|
"[%s] [%d/%d] %s → base_model=%s confidence=%s",
|
||||||
|
skill_name, processed + 1, total, model_filename,
|
||||||
|
(llm_response.get("base_model") or "?")[:50],
|
||||||
|
llm_response.get("confidence", "?"),
|
||||||
|
)
|
||||||
|
|
||||||
|
model_result = await post_processor.process(
|
||||||
|
skill_name=skill_name,
|
||||||
|
model_path=model_path,
|
||||||
|
llm_output=llm_response or {},
|
||||||
|
metadata=metadata,
|
||||||
|
readme_content=prompt_vars.get("readme_content_full", ""),
|
||||||
|
)
|
||||||
|
|
||||||
|
if model_result.get("success", True):
|
||||||
|
success_count += 1
|
||||||
|
uf = model_result.get("updated_fields", [])
|
||||||
|
if uf:
|
||||||
|
updated_models.append({"path": model_path, "updated_fields": uf})
|
||||||
|
updated_data = model_result.get("updates", {})
|
||||||
|
if "preview_url" in updated_data and updated_data["preview_url"]:
|
||||||
|
updated_data["preview_url"] = config.get_preview_static_url(
|
||||||
|
updated_data["preview_url"]
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
errors.extend(
|
||||||
|
model_result.get("errors", [model_result.get("error", "Unknown error")])
|
||||||
|
)
|
||||||
|
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("Skill %s failed for %s: %s", skill_name, model_path, exc)
|
||||||
|
errors.append(f"{model_path}: {exc}")
|
||||||
|
|
||||||
|
processed += 1
|
||||||
|
await self._emit_progress(
|
||||||
|
progress_callback, skill_name, status="processing",
|
||||||
|
total=total, processed=processed, success=success_count,
|
||||||
|
skipped=skipped_count,
|
||||||
|
current_path=model_path,
|
||||||
|
updated_data=updated_data,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = SkillResult(
|
||||||
|
success=success_count > 0,
|
||||||
|
updated_models=updated_models,
|
||||||
|
errors=errors,
|
||||||
|
summary=f"Processed {processed}/{total} models, {success_count} succeeded, {skipped_count} skipped",
|
||||||
|
)
|
||||||
|
|
||||||
|
await self._emit_progress(
|
||||||
|
progress_callback, skill_name, status="completed",
|
||||||
|
total=total, processed=processed, success=success_count,
|
||||||
|
skipped=skipped_count,
|
||||||
|
updated_models=updated_models, errors=errors, summary=result.summary,
|
||||||
|
)
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Base model grouping (keeps the prompt compact)
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _format_base_models(models: List[str]) -> str:
|
||||||
|
"""Format the base model list as a flat, one-per-line list.
|
||||||
|
|
||||||
|
Attempts to group by family consistently degraded LLM extraction
|
||||||
|
accuracy — the LLM finds individual model names harder to spot
|
||||||
|
in comma-separated groups than in a simple ``- Name`` list.
|
||||||
|
"""
|
||||||
|
return "\n".join(f"- {m}" for m in models)
|
||||||
|
|
||||||
|
async def _build_prompt_context(
|
||||||
|
self,
|
||||||
|
skill_name: str,
|
||||||
|
model_path: str,
|
||||||
|
metadata: Dict[str, Any],
|
||||||
|
registry: SkillRegistry,
|
||||||
|
llm: Any,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""Gather variables for the skill's prompt template.
|
||||||
|
|
||||||
|
Reads metadata, fetches the HF README (if applicable), lists available
|
||||||
|
base models, loads user priority tags, and returns a dict that maps to
|
||||||
|
``{{variable}}`` placeholders in ``prompt.md``.
|
||||||
|
"""
|
||||||
|
from ...metadata_ops import identify_model_type, list_base_models
|
||||||
|
from ..settings_manager import SettingsManager
|
||||||
|
|
||||||
|
context: Dict[str, Any] = {
|
||||||
|
"model_path": model_path,
|
||||||
|
"model_basename": "",
|
||||||
|
"hf_url": "",
|
||||||
|
"repo": "",
|
||||||
|
"readme_content": "",
|
||||||
|
"readme_content_full": "",
|
||||||
|
"current_metadata": {},
|
||||||
|
"base_models": [],
|
||||||
|
"priority_tags": "",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Extract model basename (filename without extension) for the LLM
|
||||||
|
# to use when locating the matching section in collection repos.
|
||||||
|
raw_basename = os.path.splitext(os.path.basename(model_path))[0]
|
||||||
|
context["model_basename"] = raw_basename or ""
|
||||||
|
|
||||||
|
context["current_metadata"] = {
|
||||||
|
"file_name": metadata.get("file_name", ""),
|
||||||
|
"base_model": metadata.get("base_model", ""),
|
||||||
|
"tags": metadata.get("tags", []),
|
||||||
|
"modelDescription": metadata.get("modelDescription", ""),
|
||||||
|
"trainedWords": metadata.get("trainedWords", []),
|
||||||
|
"sha256": (metadata.get("sha256") or "")[:16] + "..." if metadata.get("sha256") else "",
|
||||||
|
"size": metadata.get("size", 0),
|
||||||
|
}
|
||||||
|
|
||||||
|
hf_url = metadata.get("hf_url", "")
|
||||||
|
context["hf_url"] = hf_url
|
||||||
|
repo = self._extract_repo_from_url(hf_url) if hf_url else ""
|
||||||
|
context["repo"] = repo or ""
|
||||||
|
if repo:
|
||||||
|
readme = await self._fetch_readme(repo)
|
||||||
|
# Trim README to the section relevant to this model file
|
||||||
|
# (collection repos often have multiple models in one README).
|
||||||
|
if readme and raw_basename:
|
||||||
|
trimmed = extract_relevant_section(readme, raw_basename)
|
||||||
|
cleaned = clean_readme_for_llm(trimmed) if trimmed else ""
|
||||||
|
else:
|
||||||
|
cleaned = clean_readme_for_llm(readme) if readme else ""
|
||||||
|
context["readme_content"] = cleaned if cleaned else "(README not available)"
|
||||||
|
context["readme_content_full"] = readme or ""
|
||||||
|
|
||||||
|
try:
|
||||||
|
raw_models = await list_base_models()
|
||||||
|
context["base_models"] = self._format_base_models(raw_models)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.debug("Failed to list base models: %s", exc)
|
||||||
|
context["base_models"] = "</not available>"
|
||||||
|
|
||||||
|
# Determine model type and load the corresponding priority_tags
|
||||||
|
try:
|
||||||
|
model_type = await identify_model_type(model_path)
|
||||||
|
context["model_type"] = model_type
|
||||||
|
settings = SettingsManager()
|
||||||
|
priority_config = settings.get_priority_tag_config()
|
||||||
|
context["priority_tags"] = priority_config.get(model_type, "")
|
||||||
|
except Exception as exc:
|
||||||
|
logger.debug("Failed to load priority tags: %s", exc)
|
||||||
|
context["model_type"] = "lora"
|
||||||
|
context["priority_tags"] = ""
|
||||||
|
|
||||||
|
return context
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _extract_repo_from_url(hf_url: str) -> Optional[str]:
|
||||||
|
"""Extract ``user/repo`` from a HuggingFace URL."""
|
||||||
|
if not hf_url:
|
||||||
|
return None
|
||||||
|
m = re.match(r"https?://huggingface\.co/([^/]+/[^/]+)", hf_url)
|
||||||
|
return m.group(1) if m else None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def _fetch_readme(repo: str) -> str:
|
||||||
|
"""Fetch README.md from HuggingFace (tries ``main``, then ``master``)."""
|
||||||
|
async with aiohttp.ClientSession(
|
||||||
|
headers={"User-Agent": "ComfyUI-LoRA-Manager/1.0"},
|
||||||
|
timeout=aiohttp.ClientTimeout(total=30),
|
||||||
|
) as session:
|
||||||
|
for branch in ("main", "master"):
|
||||||
|
url = f"https://huggingface.co/{repo}/raw/{branch}/README.md"
|
||||||
|
try:
|
||||||
|
async with session.get(url) as resp:
|
||||||
|
if resp.status == 200:
|
||||||
|
return await resp.text()
|
||||||
|
except Exception as exc:
|
||||||
|
logger.debug("Failed to fetch README from %s: %s", url, exc)
|
||||||
|
return ""
|
||||||
|
|
||||||
|
async def _emit_progress(
|
||||||
|
self,
|
||||||
|
callback: Optional[AgentProgressReporter],
|
||||||
|
skill_name: str,
|
||||||
|
*,
|
||||||
|
status: str,
|
||||||
|
**extra: Any,
|
||||||
|
) -> None:
|
||||||
|
"""Send a progress update via WebSocket (if callback is set)."""
|
||||||
|
payload: Dict[str, Any] = {"type": "agent_progress", "skill": skill_name, "status": status}
|
||||||
|
payload.update(extra)
|
||||||
|
if callback is not None:
|
||||||
|
await callback.on_progress(payload)
|
||||||
@@ -0,0 +1,336 @@
|
|||||||
|
"""Post-processing engine for skill pipeline outputs.
|
||||||
|
|
||||||
|
The :class:`PostProcessor` takes the LLM's structured JSON output and applies
|
||||||
|
it to a model's on-disk metadata via the :mod:`~py.metadata_ops` functions.
|
||||||
|
|
||||||
|
It handles all the skill-specific business logic — conditions, transformations,
|
||||||
|
and orchestration of multiple side-effects (write metadata, download preview,
|
||||||
|
refresh cache). All actual I/O is delegated to :mod:`~py.metadata_ops`.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class PostProcessor:
|
||||||
|
"""Deterministic post-processor for skill pipeline outputs.
|
||||||
|
|
||||||
|
Usage (called by :class:`~py.services.agent.agent_service.AgentService`)::
|
||||||
|
|
||||||
|
processor = PostProcessor()
|
||||||
|
result = await processor.process(
|
||||||
|
skill_name="enrich_hf_metadata",
|
||||||
|
model_path="/path/to/model.safetensors",
|
||||||
|
llm_output={...},
|
||||||
|
metadata={...}, # from metadata_ops.read_metadata()
|
||||||
|
)
|
||||||
|
"""
|
||||||
|
|
||||||
|
async def process(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
skill_name: str,
|
||||||
|
model_path: str,
|
||||||
|
llm_output: Dict[str, Any],
|
||||||
|
metadata: Dict[str, Any],
|
||||||
|
readme_content: str = "",
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""Route *llm_output* to the correct skill post-processor.
|
||||||
|
|
||||||
|
*readme_content* is optional raw markdown content (e.g. HF README)
|
||||||
|
that is converted to HTML and stored as ``modelDescription`` for
|
||||||
|
the description tab.
|
||||||
|
|
||||||
|
Returns a dict with keys ``success`` (bool), ``updated_fields`` (list),
|
||||||
|
``preview_downloaded`` (bool), and ``errors`` (list).
|
||||||
|
"""
|
||||||
|
if skill_name == "enrich_hf_metadata":
|
||||||
|
return await self._process_enrich_hf_metadata(
|
||||||
|
model_path, llm_output, metadata, readme_content,
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"updated_fields": [],
|
||||||
|
"errors": [f"No post-processor registered for skill: {skill_name}"],
|
||||||
|
}
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# enrich_hf_metadata
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _process_enrich_hf_metadata(
|
||||||
|
self,
|
||||||
|
model_path: str,
|
||||||
|
llm_output: Dict[str, Any],
|
||||||
|
metadata: Dict[str, Any],
|
||||||
|
readme_content: str = "",
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
from ...metadata_ops import (
|
||||||
|
apply_metadata_updates,
|
||||||
|
download_preview,
|
||||||
|
refresh_cache,
|
||||||
|
)
|
||||||
|
from .skills.enrich_hf_metadata.readme_processor import (
|
||||||
|
convert_readme_to_html,
|
||||||
|
extract_gallery_images,
|
||||||
|
extract_gallery_table_images,
|
||||||
|
extract_relevant_section,
|
||||||
|
extract_simple_markdown_images,
|
||||||
|
extract_html_img_tags,
|
||||||
|
extract_repo_from_hf_url,
|
||||||
|
)
|
||||||
|
|
||||||
|
updated_fields: List[str] = []
|
||||||
|
preview_downloaded = False
|
||||||
|
|
||||||
|
# -- Determine whether this is an HF-sourced model -----------------
|
||||||
|
is_hf_model = not metadata.get("from_civitai", True)
|
||||||
|
|
||||||
|
# -- Collect updates -----------------------------------------------
|
||||||
|
updates: Dict[str, Any] = {}
|
||||||
|
|
||||||
|
# base_model
|
||||||
|
new_base = (llm_output.get("base_model") or "").strip()
|
||||||
|
current_base = metadata.get("base_model", "") or ""
|
||||||
|
if new_base and self._should_overwrite(current_base, is_hf_model):
|
||||||
|
updates["base_model"] = new_base
|
||||||
|
|
||||||
|
# trigger words → civitai.trainedWords
|
||||||
|
new_triggers = llm_output.get("trigger_words", [])
|
||||||
|
trigger_words_empty = True
|
||||||
|
if isinstance(new_triggers, list):
|
||||||
|
cleaned = [t.strip() for t in new_triggers if t.strip()]
|
||||||
|
cleaned = [t for t in cleaned if t.lower() not in ("none", "null", "n/a")]
|
||||||
|
trigger_words_empty = not cleaned
|
||||||
|
current_civitai = metadata.get("civitai") or {}
|
||||||
|
current_triggers = current_civitai.get("trainedWords") or []
|
||||||
|
if self._should_overwrite_list(current_triggers, is_hf_model):
|
||||||
|
trig_civitai = dict(current_civitai)
|
||||||
|
if "civitai" in updates and isinstance(updates["civitai"], dict):
|
||||||
|
trig_civitai.update(updates["civitai"])
|
||||||
|
trig_civitai["trainedWords"] = cleaned
|
||||||
|
updates["civitai"] = trig_civitai
|
||||||
|
|
||||||
|
# modelDescription — from raw README content (converted to HTML)
|
||||||
|
if readme_content and is_hf_model:
|
||||||
|
converted = convert_readme_to_html(readme_content)
|
||||||
|
if converted:
|
||||||
|
updates["modelDescription"] = converted
|
||||||
|
|
||||||
|
# short_description → civitai.description (for "About this version")
|
||||||
|
short_desc = (llm_output.get("short_description") or "").strip()
|
||||||
|
if short_desc and is_hf_model:
|
||||||
|
current_civitai = metadata.get("civitai") or {}
|
||||||
|
desc_civitai = dict(current_civitai)
|
||||||
|
if "civitai" in updates and isinstance(updates["civitai"], dict):
|
||||||
|
desc_civitai.update(updates["civitai"])
|
||||||
|
desc_civitai["description"] = short_desc
|
||||||
|
updates["civitai"] = desc_civitai
|
||||||
|
|
||||||
|
# gallery images → civitai.images (from YAML frontmatter widget entries
|
||||||
|
# and Sample Gallery markdown tables in the README body)
|
||||||
|
gallery_images: List[Dict[str, Any]] = []
|
||||||
|
if readme_content and is_hf_model:
|
||||||
|
hf_url = metadata.get("hf_url", "") or ""
|
||||||
|
repo = extract_repo_from_hf_url(hf_url)
|
||||||
|
if repo:
|
||||||
|
rec_w = llm_output.get("recommended_width") or 0
|
||||||
|
rec_h = llm_output.get("recommended_height") or 0
|
||||||
|
|
||||||
|
# 1. Widget images (YAML frontmatter)
|
||||||
|
gallery = extract_gallery_images(
|
||||||
|
readme_content, repo,
|
||||||
|
default_width=rec_w, default_height=rec_h,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 2. Sample Gallery table images (markdown body), deduplicated
|
||||||
|
existing_urls = {img["url"] for img in gallery if img.get("url")}
|
||||||
|
table_images = extract_gallery_table_images(
|
||||||
|
readme_content, repo,
|
||||||
|
existing_urls=existing_urls,
|
||||||
|
default_width=rec_w, default_height=rec_h,
|
||||||
|
)
|
||||||
|
existing_urls.update(img["url"] for img in table_images if img.get("url"))
|
||||||
|
|
||||||
|
# 3. Simple markdown images `` in the body
|
||||||
|
simple_images = extract_simple_markdown_images(
|
||||||
|
readme_content, repo,
|
||||||
|
existing_urls=existing_urls,
|
||||||
|
default_width=rec_w, default_height=rec_h,
|
||||||
|
)
|
||||||
|
existing_urls.update(img["url"] for img in simple_images if img.get("url"))
|
||||||
|
|
||||||
|
# 4. HTML `<img>` tags (used by many collection repos)
|
||||||
|
html_images = extract_html_img_tags(
|
||||||
|
readme_content, repo,
|
||||||
|
existing_urls=existing_urls,
|
||||||
|
default_width=rec_w, default_height=rec_h,
|
||||||
|
)
|
||||||
|
|
||||||
|
all_images = gallery + table_images + simple_images + html_images
|
||||||
|
if all_images:
|
||||||
|
gallery_images = all_images
|
||||||
|
current_civitai = metadata.get("civitai") or {}
|
||||||
|
gallery_civitai = dict(current_civitai)
|
||||||
|
if "civitai" in updates and isinstance(updates["civitai"], dict):
|
||||||
|
gallery_civitai.update(updates["civitai"])
|
||||||
|
gallery_civitai["images"] = all_images
|
||||||
|
updates["civitai"] = gallery_civitai
|
||||||
|
|
||||||
|
# tags
|
||||||
|
new_tags = llm_output.get("tags", [])
|
||||||
|
if isinstance(new_tags, list) and new_tags:
|
||||||
|
existing_tags = metadata.get("tags") or []
|
||||||
|
merged = self._merge_tags(existing_tags, new_tags)
|
||||||
|
if len(merged) > len(existing_tags) or is_hf_model:
|
||||||
|
updates["tags"] = merged
|
||||||
|
|
||||||
|
# metadata_source & llm_enriched_at (always set)
|
||||||
|
updates["metadata_source"] = "agent:enrich_hf_metadata"
|
||||||
|
updates["llm_enriched_at"] = datetime.now(timezone.utc).isoformat()
|
||||||
|
|
||||||
|
# Store LLM confidence in metadata so it's accessible for evaluation
|
||||||
|
raw_confidence = (llm_output.get("confidence") or "").strip()
|
||||||
|
if raw_confidence:
|
||||||
|
updates["_llm_confidence"] = raw_confidence
|
||||||
|
|
||||||
|
# Fallback: extract instance_prompt from YAML frontmatter when the LLM
|
||||||
|
# returned empty trigger words but the README has instance_prompt.
|
||||||
|
if trigger_words_empty:
|
||||||
|
instance_prompt = _extract_yaml_instance_prompt(readme_content)
|
||||||
|
if instance_prompt:
|
||||||
|
current_civitai = metadata.get("civitai") or {}
|
||||||
|
trig_civitai = dict(current_civitai)
|
||||||
|
if "civitai" in updates and isinstance(updates["civitai"], dict):
|
||||||
|
trig_civitai.update(updates["civitai"])
|
||||||
|
trig_civitai["trainedWords"] = [instance_prompt]
|
||||||
|
updates["civitai"] = trig_civitai
|
||||||
|
|
||||||
|
preview_remote_url = (llm_output.get("preview_url") or "").strip()
|
||||||
|
# Fallback: if the LLM couldn't find a preview image in the cleaned
|
||||||
|
# README, find the first gallery image from the *model-specific
|
||||||
|
# section* of the README (not the repo-wide first image, which
|
||||||
|
# belongs to a different model in collection repos).
|
||||||
|
if not preview_remote_url and readme_content and is_hf_model:
|
||||||
|
model_basename = os.path.splitext(os.path.basename(model_path))[0]
|
||||||
|
relevant_section = extract_relevant_section(
|
||||||
|
readme_content, model_basename,
|
||||||
|
)
|
||||||
|
if relevant_section and relevant_section != readme_content:
|
||||||
|
for img in gallery_images:
|
||||||
|
img_url = img.get("url", "")
|
||||||
|
if img_url and img_url in relevant_section:
|
||||||
|
preview_remote_url = img_url
|
||||||
|
break
|
||||||
|
# Last resort: use the first gallery image from the full README.
|
||||||
|
if not preview_remote_url and gallery_images:
|
||||||
|
preview_remote_url = gallery_images[0].get("url", "")
|
||||||
|
current_preview = metadata.get("preview_url") or ""
|
||||||
|
if preview_remote_url and not (current_preview and os.path.exists(current_preview)):
|
||||||
|
local_path = await download_preview(model_path, preview_remote_url)
|
||||||
|
if local_path:
|
||||||
|
preview_downloaded = True
|
||||||
|
updates["preview_url"] = local_path
|
||||||
|
|
||||||
|
# notes — plain-text summary of usage info from the LLM
|
||||||
|
new_notes = (llm_output.get("notes") or "").strip()
|
||||||
|
if new_notes:
|
||||||
|
updates["notes"] = new_notes
|
||||||
|
|
||||||
|
# usage_tips — JSON string (e.g. {"strength_min":0.85,"strength_max":1.4})
|
||||||
|
raw_tips = (llm_output.get("usage_tips") or "").strip()
|
||||||
|
if raw_tips and raw_tips != "{}":
|
||||||
|
try:
|
||||||
|
json.loads(raw_tips)
|
||||||
|
updates["usage_tips"] = raw_tips
|
||||||
|
except (json.JSONDecodeError, TypeError):
|
||||||
|
logger.warning(
|
||||||
|
"LLM returned invalid usage_tips JSON: %s", raw_tips[:200]
|
||||||
|
)
|
||||||
|
|
||||||
|
if updates:
|
||||||
|
updated_fields = await apply_metadata_updates(model_path, updates)
|
||||||
|
|
||||||
|
# -- Refresh scanner cache ------------------------------------------
|
||||||
|
if updated_fields or preview_downloaded:
|
||||||
|
await refresh_cache(model_path)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"success": True,
|
||||||
|
"updated_fields": updated_fields,
|
||||||
|
"preview_downloaded": preview_downloaded,
|
||||||
|
"updates": updates,
|
||||||
|
"errors": [],
|
||||||
|
}
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Helpers
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _should_overwrite(current_value: str, is_hf_model: bool) -> bool:
|
||||||
|
"""Return ``True`` when a scalar field should be overwritten."""
|
||||||
|
return is_hf_model or not current_value or current_value.lower() in (
|
||||||
|
"", "unknown",
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _should_overwrite_list(current_list: List[str], is_hf_model: bool) -> bool:
|
||||||
|
"""Return ``True`` when a list field should be overwritten."""
|
||||||
|
return is_hf_model or not current_list
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _merge_tags(existing: List[str], new: List[str]) -> List[str]:
|
||||||
|
"""Merge *new* tags into *existing*, all lowercased.
|
||||||
|
|
||||||
|
This matches the behaviour of :class:`TagUpdateService` which
|
||||||
|
normalises every tag to lowercase for case-insensitive dedup.
|
||||||
|
"""
|
||||||
|
merged: List[str] = []
|
||||||
|
seen: set = set()
|
||||||
|
for tag in list(existing) + list(new):
|
||||||
|
t = tag.strip().lower()
|
||||||
|
if t and t not in seen:
|
||||||
|
merged.append(t)
|
||||||
|
seen.add(t)
|
||||||
|
return merged
|
||||||
|
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Module-level helpers
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_yaml_instance_prompt(readme_content: str) -> str:
|
||||||
|
"""Extract ``instance_prompt`` from the YAML frontmatter of a HF README.
|
||||||
|
|
||||||
|
Returns the prompt text, or empty string if not found. Handles
|
||||||
|
``null`` / ``~`` YAML null values by returning empty string.
|
||||||
|
"""
|
||||||
|
if not readme_content or not readme_content.startswith("---"):
|
||||||
|
return ""
|
||||||
|
|
||||||
|
# Find end of frontmatter
|
||||||
|
end = readme_content.find("---", 3)
|
||||||
|
if end == -1:
|
||||||
|
return ""
|
||||||
|
frontmatter = readme_content[3:end]
|
||||||
|
|
||||||
|
for line in frontmatter.split("\n"):
|
||||||
|
line = line.strip()
|
||||||
|
m = re.match(r"^instance_prompt:\s*(.*)", line)
|
||||||
|
if m:
|
||||||
|
val = m.group(1).strip().strip('"').strip("'")
|
||||||
|
if val.lower() in ("null", "~", "none", ""):
|
||||||
|
return ""
|
||||||
|
return val
|
||||||
|
|
||||||
|
return ""
|
||||||
@@ -0,0 +1,45 @@
|
|||||||
|
"""Skill definition data structures.
|
||||||
|
|
||||||
|
Each skill is described by a :class:`SkillDefinition` that declares its
|
||||||
|
input/output schemas, whether it needs an LLM call, and what permissions
|
||||||
|
its post-processor has.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any, Dict, List, Optional, Tuple
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class SkillPermissions:
|
||||||
|
"""Declarative permission scope for a skill's post-processor.
|
||||||
|
|
||||||
|
These are auditable constraints — the :class:`AgentService` checks them
|
||||||
|
before invoking the handler. They are defense-in-depth, not a sandbox.
|
||||||
|
"""
|
||||||
|
|
||||||
|
write_metadata: bool = True
|
||||||
|
write_previews: bool = True
|
||||||
|
network_domains: Tuple[str, ...] = ()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class SkillDefinition:
|
||||||
|
"""Immutable description of an agent skill."""
|
||||||
|
|
||||||
|
name: str
|
||||||
|
title: str
|
||||||
|
description: str
|
||||||
|
llm_required: bool
|
||||||
|
input_schema: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
output_schema: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
model_type_filter: Optional[List[str]] = None
|
||||||
|
permissions: SkillPermissions = field(default_factory=SkillPermissions)
|
||||||
|
|
||||||
|
def applies_to_model_type(self, model_type: str) -> bool:
|
||||||
|
"""Return ``True`` if this skill can run on the given model type."""
|
||||||
|
|
||||||
|
if self.model_type_filter is None:
|
||||||
|
return True
|
||||||
|
return model_type in self.model_type_filter
|
||||||
@@ -0,0 +1,210 @@
|
|||||||
|
"""Discovery and loading of prompt-based skills.
|
||||||
|
|
||||||
|
Skills live in ``py/services/agent/skills/<name>/`` directories. Each
|
||||||
|
directory must contain a ``prompt.md`` file with YAML frontmatter::
|
||||||
|
|
||||||
|
---
|
||||||
|
name: my_skill
|
||||||
|
title: "My Skill"
|
||||||
|
description: "What this skill does"
|
||||||
|
llm_required: true
|
||||||
|
---
|
||||||
|
|
||||||
|
Prompt template with ``{{variable}}`` placeholders.
|
||||||
|
|
||||||
|
Legacy ``SKILL.md`` files are also supported for backward compatibility.
|
||||||
|
|
||||||
|
The registry scans the skills directory on first access and caches results.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
import yaml
|
||||||
|
|
||||||
|
from .skill_definition import SkillDefinition, SkillPermissions
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Directory where built-in skills are stored
|
||||||
|
_SKILLS_DIR = Path(__file__).parent / "skills"
|
||||||
|
|
||||||
|
#: Preferred file names for prompt definition files (tried in order).
|
||||||
|
#: ``prompt.md`` is the current convention; ``SKILL.md`` is the legacy name
|
||||||
|
#: kept for backward compatibility.
|
||||||
|
_PROMPT_FILE_NAMES: tuple[str, ...] = ("prompt.md", "SKILL.md")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Frontmatter parser
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
_FRONTMATTER_RE = re.compile(
|
||||||
|
r"^---\s*\n(.*?\n)---\s*\n?(.*)", re.DOTALL
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_skill_file(path: Path) -> tuple[dict, str]:
|
||||||
|
"""Read a prompt definition file (``prompt.md`` or legacy ``SKILL.md``) and
|
||||||
|
return (frontmatter_dict, body_text).
|
||||||
|
|
||||||
|
Raises ``ValueError`` if the file lacks valid YAML frontmatter.
|
||||||
|
"""
|
||||||
|
text = path.read_text(encoding="utf-8")
|
||||||
|
m = _FRONTMATTER_RE.match(text)
|
||||||
|
if not m:
|
||||||
|
raise ValueError(f"Missing or invalid YAML frontmatter in {path}")
|
||||||
|
frontmatter = yaml.safe_load(m.group(1))
|
||||||
|
if not isinstance(frontmatter, dict):
|
||||||
|
raise ValueError(f"Frontmatter in {path} is not a mapping")
|
||||||
|
body = m.group(2).strip()
|
||||||
|
return frontmatter, body
|
||||||
|
|
||||||
|
|
||||||
|
class SkillRegistry:
|
||||||
|
"""Discover and load agent skills from the filesystem."""
|
||||||
|
|
||||||
|
_instance: Optional["SkillRegistry"] = None
|
||||||
|
_lock: asyncio.Lock = asyncio.Lock()
|
||||||
|
|
||||||
|
def __init__(self, skills_dir: Path = _SKILLS_DIR) -> None:
|
||||||
|
self._skills_dir = skills_dir
|
||||||
|
self._skills: Dict[str, SkillDefinition] = {}
|
||||||
|
self._loaded: bool = False
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Singleton access
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def get_instance(cls) -> "SkillRegistry":
|
||||||
|
"""Return the lazily-initialised global ``SkillRegistry``."""
|
||||||
|
|
||||||
|
if cls._instance is None:
|
||||||
|
async with cls._lock:
|
||||||
|
if cls._instance is None:
|
||||||
|
registry = cls()
|
||||||
|
registry._discover()
|
||||||
|
cls._instance = registry
|
||||||
|
return cls._instance
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def reset_instance(cls) -> None:
|
||||||
|
"""Reset the cached singleton — primarily for tests."""
|
||||||
|
|
||||||
|
cls._instance = None
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Discovery
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _find_prompt_file(skill_dir: Path) -> Path | None:
|
||||||
|
"""Return the first prompt definition file that exists in *skill_dir*.
|
||||||
|
|
||||||
|
Tries ``_PROMPT_FILE_NAMES`` in order so that new conventions
|
||||||
|
(``prompt.md``) take precedence while legacy ``SKILL.md`` files
|
||||||
|
still load without changes.
|
||||||
|
"""
|
||||||
|
for name in _PROMPT_FILE_NAMES:
|
||||||
|
candidate = skill_dir / name
|
||||||
|
if candidate.exists():
|
||||||
|
return candidate
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _discover(self) -> None:
|
||||||
|
"""Scan the skills directory and load all valid skill definitions."""
|
||||||
|
|
||||||
|
self._skills.clear()
|
||||||
|
if not self._skills_dir.is_dir():
|
||||||
|
logger.warning("Skills directory does not exist: %s", self._skills_dir)
|
||||||
|
self._loaded = True
|
||||||
|
return
|
||||||
|
|
||||||
|
for entry in sorted(self._skills_dir.iterdir()):
|
||||||
|
if not entry.is_dir():
|
||||||
|
continue
|
||||||
|
prompt_file = self._find_prompt_file(entry)
|
||||||
|
if prompt_file is None:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
definition = self._load_skill_definition(prompt_file)
|
||||||
|
if definition is not None:
|
||||||
|
self._skills[definition.name] = definition
|
||||||
|
logger.debug("Loaded skill: %s", definition.name)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Failed to load skill from %s: %s", prompt_file, exc)
|
||||||
|
|
||||||
|
self._loaded = True
|
||||||
|
logger.info("Discovered %d prompt-based skills", len(self._skills))
|
||||||
|
|
||||||
|
def _load_skill_definition(self, path: Path) -> Optional[SkillDefinition]:
|
||||||
|
"""Parse a prompt definition file's frontmatter into a
|
||||||
|
:class:`SkillDefinition`."""
|
||||||
|
|
||||||
|
try:
|
||||||
|
data, _body = _parse_skill_file(path)
|
||||||
|
except (ValueError, yaml.YAMLError) as exc:
|
||||||
|
logger.warning("Failed to parse prompt file %s: %s", path, exc)
|
||||||
|
return None
|
||||||
|
|
||||||
|
if "name" not in data:
|
||||||
|
logger.warning("Prompt file %s missing required 'name' field", path)
|
||||||
|
return None
|
||||||
|
|
||||||
|
perm_data = data.get("permissions", {})
|
||||||
|
permissions = SkillPermissions(
|
||||||
|
write_metadata=perm_data.get("write_metadata", True),
|
||||||
|
write_previews=perm_data.get("write_previews", True),
|
||||||
|
network_domains=tuple(perm_data.get("network_domains", [])),
|
||||||
|
)
|
||||||
|
|
||||||
|
return SkillDefinition(
|
||||||
|
name=data["name"],
|
||||||
|
title=data.get("title", data["name"]),
|
||||||
|
description=data.get("description", ""),
|
||||||
|
llm_required=data.get("llm_required", False),
|
||||||
|
input_schema=data.get("input_schema", {}),
|
||||||
|
output_schema=data.get("output_schema", {}),
|
||||||
|
model_type_filter=data.get("model_type_filter"),
|
||||||
|
permissions=permissions,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Public API
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def list_skills(self) -> List[SkillDefinition]:
|
||||||
|
"""Return all discovered skill definitions."""
|
||||||
|
|
||||||
|
if not self._loaded:
|
||||||
|
self._discover()
|
||||||
|
return list(self._skills.values())
|
||||||
|
|
||||||
|
def get_skill(self, name: str) -> Optional[SkillDefinition]:
|
||||||
|
"""Return the skill definition for ``name``, or ``None`` if not found."""
|
||||||
|
|
||||||
|
if not self._loaded:
|
||||||
|
self._discover()
|
||||||
|
return self._skills.get(name)
|
||||||
|
|
||||||
|
def load_prompt(self, name: str) -> str:
|
||||||
|
"""Load and return the prompt template body for the named skill."""
|
||||||
|
|
||||||
|
skill_dir = self._skills_dir / name
|
||||||
|
skill_path = self._find_prompt_file(skill_dir)
|
||||||
|
if skill_path is None:
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"Prompt file not found for skill '{name}' in {skill_dir} "
|
||||||
|
f"(tried {list(_PROMPT_FILE_NAMES)})"
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
_frontmatter, body = _parse_skill_file(skill_path)
|
||||||
|
return body
|
||||||
|
except (ValueError, yaml.YAMLError) as exc:
|
||||||
|
raise ValueError(f"Failed to parse prompt from {skill_path}: {exc}") from exc
|
||||||
@@ -0,0 +1,165 @@
|
|||||||
|
---
|
||||||
|
name: enrich_hf_metadata
|
||||||
|
title: "Enrich Metadata from HuggingFace"
|
||||||
|
description: >
|
||||||
|
Parse the HuggingFace model card via LLM to extract description, trigger
|
||||||
|
words, base model, tags, and preview image URL.
|
||||||
|
llm_required: true
|
||||||
|
---
|
||||||
|
|
||||||
|
You are an expert assistant for AI image generation models. Your task is to extract structured metadata from a HuggingFace model card (README.md).
|
||||||
|
|
||||||
|
## Model Information
|
||||||
|
|
||||||
|
- **Repository**: {{hf_url}}
|
||||||
|
- **Model file path**: {{model_path}}
|
||||||
|
- **Model filename**: {{model_basename}}
|
||||||
|
- **Repository ID**: {{repo}}
|
||||||
|
|
||||||
|
## Current Metadata (may be incomplete)
|
||||||
|
|
||||||
|
```json
|
||||||
|
{{current_metadata}}
|
||||||
|
```
|
||||||
|
|
||||||
|
## User Priority Tags Reference
|
||||||
|
|
||||||
|
The user has configured the following list of **meaningful tag categories** for this model type (`{{model_type}}`):
|
||||||
|
|
||||||
|
```
|
||||||
|
{{priority_tags}}
|
||||||
|
```
|
||||||
|
|
||||||
|
These are the subjects, styles, and concepts the user considers useful for categorization. Use this list as a **reference** when evaluating tags (see the **tags** section below).
|
||||||
|
|
||||||
|
## Available Base Models
|
||||||
|
|
||||||
|
The following base models are currently valid in this system. Use the EXACT
|
||||||
|
name listed — do not invent aliases or modify variant suffixes.
|
||||||
|
|
||||||
|
{{base_models}}
|
||||||
|
|
||||||
|
## HuggingFace README Content
|
||||||
|
|
||||||
|
```
|
||||||
|
{{readme_content}}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Extraction Instructions
|
||||||
|
|
||||||
|
Extract the following information from the README content above:
|
||||||
|
|
||||||
|
### base_model
|
||||||
|
The base model this model was trained on. Use EXACTLY one of the names from the **Available Base Models** list above. Do not invent new names or use aliases.
|
||||||
|
|
||||||
|
Check the YAML frontmatter for ``base_model:`` first. If the frontmatter has no ``base_model:``, look at the **model filename** (``{{model_basename}}``), YAML ``tags:``, README title and first paragraph for clues — the base model family is often embedded in the name
|
||||||
|
|
||||||
|
### trigger_words
|
||||||
|
The trigger words or activation prompts needed to use this LoRA. Look for:
|
||||||
|
- `instance_prompt:` in the YAML frontmatter
|
||||||
|
- Phrases like "trigger word:", "trigger:", "use this prompt:", "activation prompt:"
|
||||||
|
- In collection repos: the trigger section **specific to this model file** (look near matching download links or anchor IDs)
|
||||||
|
- Example prompts at the start (usually the first word or phrase before any description)
|
||||||
|
Return as an array of strings. If none found, return an empty array `[]`. **Never** return `["None"]` or any placeholder value — a truly empty list means no trigger words exist.
|
||||||
|
|
||||||
|
### short_description
|
||||||
|
A concise 1-2 sentence summary of what this model does. Extract from the "Model description" section or the first paragraph. For collection repos, focus on the **specific model version** matching `{{model_basename}}`, not the repo as a whole. Return empty string if the README is too minimal.
|
||||||
|
|
||||||
|
### tags
|
||||||
|
3-8 relevant tags for categorizing this model. **Quality over quantity.**
|
||||||
|
|
||||||
|
Sources to consider:
|
||||||
|
- The YAML frontmatter `tags:` list (filter out technical ones — see below)
|
||||||
|
- The subject, style, character, or concept the model represents
|
||||||
|
- The model filename itself may give clues (e.g. "pokemon", "anime", "pixelart")
|
||||||
|
|
||||||
|
**Critical filtering rules — apply them strictly:**
|
||||||
|
|
||||||
|
1. **Exclude technical/generic tags.** Reject any tag that describes the model's **training methodology, framework, architecture, or modality** rather than its content. Examples to exclude: `text-to-image`, `diffusers`, `lora`, `dreambooth`, `diffusers-training`, `flux`, `sdxl`, `checkpoint`, `pytorch`, `safetensors`, `fine-tuning`, `stable-diffusion`, and any variant of these.
|
||||||
|
|
||||||
|
2. **Cross-reference against the priority_tags reference.** Only include a tag if it meaningfully describes what the model actually creates (subject, style, character type) and is semantically close to one of the priority_tags. If none of the README's tags match meaningful categories, prefer returning a smaller set or an empty array over including low-value tags.
|
||||||
|
|
||||||
|
3. **All lowercase, no spaces, no hyphens** (use single words like `"photorealistic"`, `"anime"`, `"character"`).
|
||||||
|
|
||||||
|
Return empty array if no meaningful content tags remain after filtering.
|
||||||
|
|
||||||
|
### recommended_width, recommended_height
|
||||||
|
The recommended image generation resolution for this model, in pixels. Look for sections like "Best Dimensions", "Recommended size", "Suggested resolution", or similar phrasing in the README. Prefer the explicitly marked "Best" or default resolution. If the table/list has multiple entries (e.g. "768 x 1024 (Best)" and "1024 x 1024 (Default)"), use the one marked "Best". Return integers. If no resolution can be determined, return 0 for both.
|
||||||
|
|
||||||
|
### preview_url
|
||||||
|
The URL of the most suitable preview image from the README. Look for:
|
||||||
|
- Image tags near the section matching the model filename (`{{model_basename}}`)
|
||||||
|
- The YAML frontmatter `widget:` section (which often has `output.url` fields)
|
||||||
|
- In collection repos: the sample images listed **under the section** for this specific model version
|
||||||
|
- Generic `` in the body
|
||||||
|
Choose the first image that appears to be a generation example (not a logo or diagram). Construct the absolute URL as `https://huggingface.co/{{repo}}/resolve/main/{filename}`. If no suitable image is found, return an empty string.
|
||||||
|
|
||||||
|
### notes
|
||||||
|
A plain-text summary of the model card's key practical usage information. Combine trigger words, style modifiers, recommended parameters (steps, CFG, resolution, sampler), and any setup tips into a readable paragraph. For collection repos, focus on the **specific model version** matching `{{model_basename}}`. Return empty string if the README has no useful usage info.
|
||||||
|
|
||||||
|
### usage_tips
|
||||||
|
A JSON string with structured usage recommendations. Extract from the README any explicit ranges or recommended values (e.g. "Set LoRA strength: **0.85 - 1.4**", "CLIP strength: 0.5"). Possible fields (include only those you can determine):
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"strength_min": 0.85,
|
||||||
|
"strength_max": 1.4,
|
||||||
|
"strength_range": "0.85-1.4",
|
||||||
|
"strength": 0.6,
|
||||||
|
"clip_strength": 0.5,
|
||||||
|
"clip_skip": 2
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Return the JSON string (e.g. `'{"strength_min":0.85,"strength_max":1.4}'`). Return `"{}"` if nothing useful is found.
|
||||||
|
|
||||||
|
### confidence
|
||||||
|
Your confidence level in the extracted data:
|
||||||
|
- "high" — most fields were explicitly stated in the README
|
||||||
|
- "medium" — some fields were inferred from context
|
||||||
|
- "low" — most fields are guesses based on limited information
|
||||||
|
|
||||||
|
## Important: Handling Collection Repos (multiple model files)
|
||||||
|
|
||||||
|
Many HuggingFace repos contain **multiple model files** in a single repository
|
||||||
|
(e.g. a "LoRA collection" with different styles/characters in separate files).
|
||||||
|
|
||||||
|
The model file currently being enriched is: **`{{model_basename}}`**
|
||||||
|
|
||||||
|
To find the correct section in the README:
|
||||||
|
|
||||||
|
1. **Search for download links** containing the filename — the surrounding paragraph is your section.
|
||||||
|
2. **Search for anchor IDs** (`<a id="...">`) or section headings whose text matches words from the filename.
|
||||||
|
3. **Search for HTML headings** (`<h1>`, `<h2>`, `<span>`) containing parts of the filename.
|
||||||
|
4. If no match is found, use the full README as usual — the model may be the only one in the repo.
|
||||||
|
|
||||||
|
When a matching section IS found, prefer metadata from that section.
|
||||||
|
When no section matches (e.g. single-model repos or repos without per-file sections),
|
||||||
|
extract metadata from the full README normally. Do not return empty data just
|
||||||
|
because the filename doesn't appear in the README.
|
||||||
|
|
||||||
|
## Output Format
|
||||||
|
|
||||||
|
Return ONLY a JSON object with exactly these fields (no markdown fences, no extra text):
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"model_path": "{{model_path}}",
|
||||||
|
"base_model": "<canonical name or empty string>",
|
||||||
|
"trigger_words": ["<word1>", "<word2>"],
|
||||||
|
"short_description": "<1-2 sentence summary>",
|
||||||
|
"tags": ["<tag1>", "<tag2>"],
|
||||||
|
"recommended_width": 768,
|
||||||
|
"recommended_height": 1024,
|
||||||
|
"preview_url": "<image URL or empty string>",
|
||||||
|
"notes": "<plain-text usage summary or empty string>",
|
||||||
|
"usage_tips": "<JSON string like '{\"strength_min\":0.85,\"strength_max\":1.4}' or '{}'>",
|
||||||
|
"confidence": "<high|medium|low>"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Important:
|
||||||
|
- Only include the JSON object, no other text
|
||||||
|
- If a field cannot be determined, use an empty string or empty array
|
||||||
|
- Do not fabricate information not supported by the README
|
||||||
|
- Never use placeholder values like `"None"` or `"unknown"` for missing data — use empty string or empty array
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -1,7 +1,7 @@
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
import asyncio
|
import asyncio
|
||||||
import re
|
import re
|
||||||
from typing import Any, Dict, List, Optional, Type, TYPE_CHECKING
|
from typing import Any, Dict, List, Optional, Type, Union, TYPE_CHECKING
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
@@ -109,12 +109,15 @@ class BaseModelService(ABC):
|
|||||||
if civitai_model_id is not None:
|
if civitai_model_id is not None:
|
||||||
sorted_data = [
|
sorted_data = [
|
||||||
item for item in sorted_data
|
item for item in sorted_data
|
||||||
if self._extract_model_id(item) == civitai_model_id
|
if self._extract_group_key(item) == civitai_model_id
|
||||||
]
|
]
|
||||||
# VLM mode: always sort by version ID descending (newest version first),
|
# VLM mode: always sort by version ID descending (newest version first),
|
||||||
# regardless of the current sort_by preference.
|
# regardless of the current sort_by preference.
|
||||||
|
# Fall back to modified timestamp for non-CivitAI sources.
|
||||||
sorted_data.sort(
|
sorted_data.sort(
|
||||||
key=lambda x: self._extract_version_id(x) or 0,
|
key=lambda x: self._extract_version_id(x)
|
||||||
|
or x.get("modified", 0)
|
||||||
|
or 0,
|
||||||
reverse=True,
|
reverse=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -129,18 +132,21 @@ class BaseModelService(ABC):
|
|||||||
ufs = self.settings.get("version_grouping", "same_base")
|
ufs = self.settings.get("version_grouping", "same_base")
|
||||||
group_by_base = ufs == "same_base"
|
group_by_base = ufs == "same_base"
|
||||||
|
|
||||||
dedup_map = {} # (modelId [,base_model]) -> (item, version_id)
|
dedup_map = {} # (modelId [,base_model]) -> (item, version_or_modified)
|
||||||
version_counter = {} # same-key -> count
|
version_counter = {} # same-key -> count
|
||||||
standalone = []
|
standalone = []
|
||||||
for item in sorted_data:
|
for item in sorted_data:
|
||||||
mid = self._extract_model_id(item)
|
mid = self._extract_group_key(item)
|
||||||
if mid is None:
|
if mid is None:
|
||||||
standalone.append(item)
|
standalone.append(item)
|
||||||
continue
|
continue
|
||||||
key = (mid, item.get("base_model") or "") if group_by_base else mid
|
key = (mid, item.get("base_model") or "") if group_by_base else mid
|
||||||
# Count all versions per key
|
# Count all versions per key
|
||||||
version_counter[key] = version_counter.get(key, 0) + 1
|
version_counter[key] = version_counter.get(key, 0) + 1
|
||||||
vid = self._extract_version_id(item) or 0
|
# Prefer CivitAI version_id; fall back to modified timestamp
|
||||||
|
vid = self._extract_version_id(item)
|
||||||
|
if vid is None:
|
||||||
|
vid = item.get("modified", 0) or 0
|
||||||
if key not in dedup_map or vid > dedup_map[key][1]:
|
if key not in dedup_map or vid > dedup_map[key][1]:
|
||||||
dedup_map[key] = (item, vid)
|
dedup_map[key] = (item, vid)
|
||||||
# Attach version_count to each surviving grouped item (shallow copy
|
# Attach version_count to each surviving grouped item (shallow copy
|
||||||
@@ -174,16 +180,19 @@ class BaseModelService(ABC):
|
|||||||
model_groups: Dict[Any, List[Dict]] = {}
|
model_groups: Dict[Any, List[Dict]] = {}
|
||||||
ungrouped_standalone: List[Dict] = []
|
ungrouped_standalone: List[Dict] = []
|
||||||
for item in sorted_data:
|
for item in sorted_data:
|
||||||
mid = self._extract_model_id(item)
|
mid = self._extract_group_key(item)
|
||||||
if mid is None:
|
if mid is None:
|
||||||
ungrouped_standalone.append(item)
|
ungrouped_standalone.append(item)
|
||||||
continue
|
continue
|
||||||
key = (mid, item.get("base_model") or "") if group_by_base else mid
|
key = (mid, item.get("base_model") or "") if group_by_base else mid
|
||||||
model_groups.setdefault(key, []).append(item)
|
model_groups.setdefault(key, []).append(item)
|
||||||
# Sort versions within each group by version id descending
|
# Sort versions within each group by version id (descending);
|
||||||
|
# fall back to modified timestamp for non-CivitAI sources.
|
||||||
for items in model_groups.values():
|
for items in model_groups.values():
|
||||||
items.sort(
|
items.sort(
|
||||||
key=lambda x: self._extract_version_id(x) or 0,
|
key=lambda x: self._extract_version_id(x)
|
||||||
|
or x.get("modified", 0)
|
||||||
|
or 0,
|
||||||
reverse=True,
|
reverse=True,
|
||||||
)
|
)
|
||||||
# Sort groups by version count
|
# Sort groups by version count
|
||||||
@@ -697,6 +706,33 @@ class BaseModelService(ABC):
|
|||||||
|
|
||||||
return annotated
|
return annotated
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _extract_hf_group_key(item: Dict) -> 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) -> 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
|
@staticmethod
|
||||||
def _extract_model_id(item: Dict) -> Optional[int]:
|
def _extract_model_id(item: Dict) -> Optional[int]:
|
||||||
civitai = item.get("civitai") if isinstance(item, dict) else None
|
civitai = item.get("civitai") if isinstance(item, dict) else None
|
||||||
@@ -804,6 +840,12 @@ class BaseModelService(ABC):
|
|||||||
"""Get top tags sorted by frequency"""
|
"""Get top tags sorted by frequency"""
|
||||||
return await self.scanner.get_top_tags(limit)
|
return await self.scanner.get_top_tags(limit)
|
||||||
|
|
||||||
|
async def search_tags(
|
||||||
|
self, query: str, limit: int = 50
|
||||||
|
) -> List[Dict]:
|
||||||
|
"""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]:
|
async def get_base_models(self, limit: int = 20) -> List[Dict]:
|
||||||
"""Get base models sorted by frequency"""
|
"""Get base models sorted by frequency"""
|
||||||
return await self.scanner.get_base_models(limit)
|
return await self.scanner.get_base_models(limit)
|
||||||
@@ -955,13 +997,21 @@ class BaseModelService(ABC):
|
|||||||
|
|
||||||
return unified_tree
|
return unified_tree
|
||||||
|
|
||||||
async def get_model_notes(self, model_name: str) -> Optional[str]:
|
async def get_model_notes(self, model_name: str) -> Optional[dict]:
|
||||||
"""Get notes for a specific model file"""
|
"""Get notes and file_path for a specific model file.
|
||||||
|
|
||||||
|
Supports both simple names (``OWSMianne_ANIMA_V1``) and full-path
|
||||||
|
syntax (``Anima/character/OWSMianne_ANIMA_V1``).
|
||||||
|
"""
|
||||||
cache = await self.scanner.get_cached_data()
|
cache = await self.scanner.get_cached_data()
|
||||||
|
|
||||||
for model in cache.raw_data:
|
for model in cache.raw_data:
|
||||||
if model["file_name"] == model_name:
|
file_name = model.get("file_name", "")
|
||||||
return model.get("notes", "")
|
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
|
return None
|
||||||
|
|
||||||
@@ -1084,6 +1134,11 @@ class BaseModelService(ABC):
|
|||||||
|
|
||||||
Listing/search endpoints return lightweight cache entries; this method performs
|
Listing/search endpoints return lightweight cache entries; this method performs
|
||||||
a lazy read of the on-disk metadata snapshot when callers need full detail.
|
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(
|
metadata, should_skip = await MetadataManager.load_metadata(
|
||||||
file_path, self.metadata_class
|
file_path, self.metadata_class
|
||||||
@@ -1101,6 +1156,19 @@ class BaseModelService(ABC):
|
|||||||
MetadataManager.save_metadata(file_path, metadata)
|
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", {}))
|
return self.filter_civitai_data(metadata.to_dict().get("civitai", {}))
|
||||||
|
|
||||||
async def get_model_description(self, file_path: str) -> Optional[str]:
|
async def get_model_description(self, file_path: str) -> Optional[str]:
|
||||||
|
|||||||
@@ -114,6 +114,13 @@ class CheckpointScanner(ModelScanner):
|
|||||||
and metadata.hash_status == "completed"
|
and metadata.hash_status == "completed"
|
||||||
and metadata.sha256
|
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)
|
||||||
return metadata.sha256
|
return metadata.sha256
|
||||||
|
|
||||||
async with self._hash_calculation_lock:
|
async with self._hash_calculation_lock:
|
||||||
@@ -125,6 +132,7 @@ class CheckpointScanner(ModelScanner):
|
|||||||
and metadata.hash_status == "completed"
|
and metadata.hash_status == "completed"
|
||||||
and metadata.sha256
|
and metadata.sha256
|
||||||
):
|
):
|
||||||
|
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
|
||||||
return metadata.sha256
|
return metadata.sha256
|
||||||
|
|
||||||
task = self._hash_calculation_tasks.get(real_path)
|
task = self._hash_calculation_tasks.get(real_path)
|
||||||
@@ -175,6 +183,9 @@ class CheckpointScanner(ModelScanner):
|
|||||||
|
|
||||||
# Check if hash is already calculated
|
# Check if hash is already calculated
|
||||||
if metadata.hash_status == "completed" and metadata.sha256:
|
if metadata.hash_status == "completed" and metadata.sha256:
|
||||||
|
# Populate the in-memory hash index even for pre-computed
|
||||||
|
# hashes, mirroring the fix in calculate_hash_for_model.
|
||||||
|
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
|
||||||
return metadata.sha256
|
return metadata.sha256
|
||||||
|
|
||||||
# Update status to calculating
|
# Update status to calculating
|
||||||
@@ -193,6 +204,20 @@ class CheckpointScanner(ModelScanner):
|
|||||||
# Update hash index
|
# Update hash index
|
||||||
self._hash_index.add_entry(sha256.lower(), file_path)
|
self._hash_index.add_entry(sha256.lower(), file_path)
|
||||||
|
|
||||||
|
# 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"
|
||||||
|
break
|
||||||
|
|
||||||
logger.info(f"Hash calculated for checkpoint: {file_path}")
|
logger.info(f"Hash calculated for checkpoint: {file_path}")
|
||||||
return sha256
|
return sha256
|
||||||
|
|
||||||
|
|||||||
@@ -304,6 +304,20 @@ class CivArchiveClient:
|
|||||||
version_id = file_data.get("model_version_id") or file_data.get("modelVersionId")
|
version_id = file_data.get("model_version_id") or file_data.get("modelVersionId")
|
||||||
if model_id is None or version_id is None:
|
if model_id is None or version_id is None:
|
||||||
continue
|
continue
|
||||||
|
# CivitAI / CivArchive model IDs are small integers (typically ≤ 7
|
||||||
|
# digits). Reject suspiciously large values that indicate the API
|
||||||
|
# returned a malformed payload (e.g. a hash reinterpreted as an ID)
|
||||||
|
# to avoid pointless HTTP 500 errors from CivArchive.
|
||||||
|
_MAX_VALID_CIVITAI_ID = 100_000_000
|
||||||
|
try:
|
||||||
|
if int(model_id) >= _MAX_VALID_CIVITAI_ID or int(version_id) >= _MAX_VALID_CIVITAI_ID:
|
||||||
|
logger.debug(
|
||||||
|
"Skipping implausible CivArchive model_id=%s / version_id=%s",
|
||||||
|
model_id, version_id,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
continue
|
||||||
resolved = await self.get_model_version(model_id, version_id)
|
resolved = await self.get_model_version(model_id, version_id)
|
||||||
if resolved:
|
if resolved:
|
||||||
return resolved
|
return resolved
|
||||||
|
|||||||
@@ -213,6 +213,18 @@ class CivitaiBaseModelService:
|
|||||||
"wan video 2.2 i2v-a14b": "WAN",
|
"wan video 2.2 i2v-a14b": "WAN",
|
||||||
"wan video 2.5 t2v": "WAN",
|
"wan video 2.5 t2v": "WAN",
|
||||||
"wan video 2.5 i2v": "WAN",
|
"wan video 2.5 i2v": "WAN",
|
||||||
|
"wan video 2.7": "WAN",
|
||||||
|
"wan image 2.7": "WI27",
|
||||||
|
"ace audio": "ACE",
|
||||||
|
"boogu": "BOOG",
|
||||||
|
"grok": "GROK",
|
||||||
|
"happyhorse": "HAPP",
|
||||||
|
"hidream-o1": "HIO1",
|
||||||
|
"lens": "LENS",
|
||||||
|
"mai": "MAI",
|
||||||
|
"upscaler": "UPSC",
|
||||||
|
"ideogram 4.0": "ID40",
|
||||||
|
"qwen 2": "QWN2",
|
||||||
}
|
}
|
||||||
|
|
||||||
if lower_name in special_cases:
|
if lower_name in special_cases:
|
||||||
@@ -392,6 +404,7 @@ class CivitaiBaseModelService:
|
|||||||
"LTXV2",
|
"LTXV2",
|
||||||
"LTXV 2.3",
|
"LTXV 2.3",
|
||||||
"CogVideoX",
|
"CogVideoX",
|
||||||
|
"HappyHorse",
|
||||||
"Mochi",
|
"Mochi",
|
||||||
"Hunyuan Video",
|
"Hunyuan Video",
|
||||||
"Wan Video",
|
"Wan Video",
|
||||||
@@ -404,15 +417,25 @@ class CivitaiBaseModelService:
|
|||||||
"Wan Video 2.2 I2V-A14B",
|
"Wan Video 2.2 I2V-A14B",
|
||||||
"Wan Video 2.5 T2V",
|
"Wan Video 2.5 T2V",
|
||||||
"Wan Video 2.5 I2V",
|
"Wan Video 2.5 I2V",
|
||||||
|
"Wan Image 2.7",
|
||||||
|
"Wan Video 2.7",
|
||||||
],
|
],
|
||||||
"Other Models": [
|
"Other Models": [
|
||||||
|
"ACE Audio",
|
||||||
"Illustrious",
|
"Illustrious",
|
||||||
"Pony",
|
"Pony",
|
||||||
"Pony V7",
|
"Pony V7",
|
||||||
|
"Boogu",
|
||||||
"HiDream",
|
"HiDream",
|
||||||
|
"HiDream-O1",
|
||||||
|
"Ideogram 4.0",
|
||||||
"Qwen",
|
"Qwen",
|
||||||
|
"Qwen 2",
|
||||||
"AuraFlow",
|
"AuraFlow",
|
||||||
"Chroma",
|
"Chroma",
|
||||||
|
"Grok",
|
||||||
|
"Lens",
|
||||||
|
"MAI",
|
||||||
"ZImageTurbo",
|
"ZImageTurbo",
|
||||||
"ZImageBase",
|
"ZImageBase",
|
||||||
"PixArt a",
|
"PixArt a",
|
||||||
@@ -426,6 +449,7 @@ class CivitaiBaseModelService:
|
|||||||
"Ernie Turbo",
|
"Ernie Turbo",
|
||||||
"Nucleus",
|
"Nucleus",
|
||||||
"Krea 2",
|
"Krea 2",
|
||||||
|
"Upscaler",
|
||||||
],
|
],
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -230,6 +230,12 @@ class DownloadManager:
|
|||||||
Returns:
|
Returns:
|
||||||
Dict with download result
|
Dict with download result
|
||||||
"""
|
"""
|
||||||
|
logger.debug(
|
||||||
|
"[download] download_from_civitai called: model_id=%s, model_version_id=%s, "
|
||||||
|
"source=%s, file_params=%s",
|
||||||
|
model_id, model_version_id, source, file_params,
|
||||||
|
)
|
||||||
|
|
||||||
# Validate that at least one identifier is provided
|
# Validate that at least one identifier is provided
|
||||||
if not model_id and not model_version_id:
|
if not model_id and not model_version_id:
|
||||||
return {
|
return {
|
||||||
@@ -250,6 +256,7 @@ class DownloadManager:
|
|||||||
"source": source,
|
"source": source,
|
||||||
"file_params": copy.deepcopy(file_params) if file_params is not None else None,
|
"file_params": copy.deepcopy(file_params) if file_params is not None else None,
|
||||||
"progress": 0,
|
"progress": 0,
|
||||||
|
|
||||||
"status": "queued",
|
"status": "queued",
|
||||||
"transfer_backend": self._get_model_download_backend(),
|
"transfer_backend": self._get_model_download_backend(),
|
||||||
"bytes_downloaded": 0,
|
"bytes_downloaded": 0,
|
||||||
@@ -289,8 +296,8 @@ class DownloadManager:
|
|||||||
return result
|
return result
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
return {
|
return {
|
||||||
"success": False,
|
"success": True,
|
||||||
"error": "Download was cancelled",
|
"cancelled": True,
|
||||||
"download_id": task_id,
|
"download_id": task_id,
|
||||||
}
|
}
|
||||||
finally:
|
finally:
|
||||||
@@ -675,7 +682,10 @@ class DownloadManager:
|
|||||||
u for u in download_urls if not u.startswith(CIVITAI_DOWNLOAD_URL_PREFIXES)
|
u for u in download_urls if not u.startswith(CIVITAI_DOWNLOAD_URL_PREFIXES)
|
||||||
]
|
]
|
||||||
download_urls = non_civitai_urls + civitai_urls
|
download_urls = non_civitai_urls + civitai_urls
|
||||||
else:
|
|
||||||
|
# Fallback: when mirrors is empty or all mirrors have been deleted,
|
||||||
|
# use the file's downloadUrl directly (e.g. CivitAI download endpoint).
|
||||||
|
if not download_urls:
|
||||||
download_url = file_info.get("downloadUrl")
|
download_url = file_info.get("downloadUrl")
|
||||||
if download_url:
|
if download_url:
|
||||||
download_urls.append(normalize_civitai_download_url(download_url))
|
download_urls.append(normalize_civitai_download_url(download_url))
|
||||||
@@ -1379,7 +1389,17 @@ class DownloadManager:
|
|||||||
|
|
||||||
# Update save directory with relative path if provided
|
# Update save directory with relative path if provided
|
||||||
if relative_path:
|
if relative_path:
|
||||||
|
base_save_dir = save_dir
|
||||||
save_dir = os.path.join(save_dir, relative_path)
|
save_dir = os.path.join(save_dir, relative_path)
|
||||||
|
# Security: validate path containment after joining
|
||||||
|
resolved_dir = os.path.abspath(os.path.normpath(save_dir))
|
||||||
|
base_dir = os.path.abspath(os.path.normpath(base_save_dir))
|
||||||
|
if not resolved_dir.startswith(base_dir + os.sep) and resolved_dir != base_dir:
|
||||||
|
logger.warning(
|
||||||
|
"Path traversal detected: %s escapes %s",
|
||||||
|
resolved_dir, base_dir,
|
||||||
|
)
|
||||||
|
return {"success": False, "error": "Download path is outside allowed directory"}
|
||||||
# Create directory if it doesn't exist
|
# Create directory if it doesn't exist
|
||||||
os.makedirs(save_dir, exist_ok=True)
|
os.makedirs(save_dir, exist_ok=True)
|
||||||
|
|
||||||
@@ -1421,14 +1441,35 @@ class DownloadManager:
|
|||||||
|
|
||||||
# If file_params is provided, try to find matching file
|
# If file_params is provided, try to find matching file
|
||||||
if file_params and model_version_id:
|
if file_params and model_version_id:
|
||||||
|
target_file_id = file_params.get("id")
|
||||||
target_type = file_params.get("type", "Model")
|
target_type = file_params.get("type", "Model")
|
||||||
target_format = file_params.get("format", "SafeTensor")
|
target_format = file_params.get("format")
|
||||||
target_size = file_params.get("size", "full")
|
target_size = file_params.get("size")
|
||||||
target_fp = file_params.get("fp")
|
target_fp = file_params.get("fp")
|
||||||
is_primary = file_params.get("isPrimary", False)
|
is_primary = file_params.get("isPrimary", False)
|
||||||
|
|
||||||
if is_primary:
|
logger.debug(
|
||||||
# Find primary file
|
"[download] file_params received: id=%s, type=%s, format=%s, size=%s, fp=%s, isPrimary=%s, "
|
||||||
|
"model_version_id=%s, total_files=%d",
|
||||||
|
target_file_id, target_type, target_format, target_size, target_fp, is_primary,
|
||||||
|
model_version_id, len(files),
|
||||||
|
)
|
||||||
|
|
||||||
|
if target_file_id:
|
||||||
|
target_id_str = str(target_file_id)
|
||||||
|
for f in files:
|
||||||
|
f_id = f.get("id")
|
||||||
|
if str(f_id) == target_id_str:
|
||||||
|
file_info = f
|
||||||
|
logger.debug(
|
||||||
|
"[download] MATCH by ID: id=%s name='%s'",
|
||||||
|
f_id, f.get("name"),
|
||||||
|
)
|
||||||
|
break
|
||||||
|
if not file_info:
|
||||||
|
logger.debug("[download] No file found with id=%s", target_file_id)
|
||||||
|
|
||||||
|
elif is_primary:
|
||||||
file_info = next(
|
file_info = next(
|
||||||
(
|
(
|
||||||
f
|
f
|
||||||
@@ -1439,28 +1480,41 @@ class DownloadManager:
|
|||||||
None,
|
None,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# Match by metadata
|
# Lenient metadata match: only compare fields present on both sides
|
||||||
for f in files:
|
for f in files:
|
||||||
f_type = f.get("type", "")
|
f_type = f.get("type", "")
|
||||||
f_meta = f.get("metadata", {})
|
|
||||||
|
|
||||||
# Check type match
|
|
||||||
if f_type != target_type:
|
if f_type != target_type:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# Check metadata match
|
f_meta = f.get("metadata", {})
|
||||||
if f_meta.get("format") != target_format:
|
f_format = f_meta.get("format") or f.get("format")
|
||||||
|
f_size = f_meta.get("size") or f.get("size")
|
||||||
|
f_fp = f_meta.get("fp") or f.get("fp")
|
||||||
|
|
||||||
|
if target_format and f_format != target_format:
|
||||||
continue
|
continue
|
||||||
if f_meta.get("size") != target_size:
|
if target_size and f_size and f_size != target_size:
|
||||||
continue
|
continue
|
||||||
if target_fp and f_meta.get("fp") != target_fp:
|
if target_fp and f_fp and f_fp != target_fp:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
file_info = f
|
file_info = f
|
||||||
break
|
break
|
||||||
|
|
||||||
|
if not file_info:
|
||||||
|
logger.debug(
|
||||||
|
"[download] No match found via file_params — falling back to primary file lookup",
|
||||||
|
)
|
||||||
|
elif not file_params:
|
||||||
|
logger.debug(
|
||||||
|
"[download] No file_params provided (null/None) — will use primary file lookup. "
|
||||||
|
"model_version_id=%s, total_files=%d",
|
||||||
|
model_version_id, len(files),
|
||||||
|
)
|
||||||
|
|
||||||
# Fallback to primary file if no match found
|
# Fallback to primary file if no match found
|
||||||
if not file_info:
|
if not file_info:
|
||||||
|
logger.debug("[download] Looking for primary file as fallback")
|
||||||
file_info = next(
|
file_info = next(
|
||||||
(
|
(
|
||||||
f
|
f
|
||||||
@@ -1469,38 +1523,18 @@ class DownloadManager:
|
|||||||
),
|
),
|
||||||
None,
|
None,
|
||||||
)
|
)
|
||||||
|
if file_info:
|
||||||
|
logger.debug(
|
||||||
|
"[download] Fallback primary file selected: id=%s, name=%s",
|
||||||
|
file_info.get("id"), file_info.get("name"),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.debug("[download] No primary file found in fallback lookup")
|
||||||
|
|
||||||
if not file_info:
|
if not file_info:
|
||||||
return {"success": False, "error": "No suitable file found in metadata"}
|
return {"success": False, "error": "No suitable file found in metadata"}
|
||||||
mirrors = file_info.get("mirrors") or []
|
|
||||||
download_urls = []
|
|
||||||
if mirrors:
|
|
||||||
for mirror in mirrors:
|
|
||||||
if mirror.get("deletedAt") is None and mirror.get("url"):
|
|
||||||
download_urls.append(
|
|
||||||
normalize_civitai_download_url(mirror["url"])
|
|
||||||
)
|
|
||||||
|
|
||||||
# When source is 'civarchive', prioritize non-Civitai URLs
|
download_urls = self._build_download_urls_from_file_info(file_info, source=source)
|
||||||
# This avoids failed downloads from deleted Civitai models
|
|
||||||
if source == "civarchive" and len(download_urls) > 1:
|
|
||||||
civitai_urls = [
|
|
||||||
u
|
|
||||||
for u in download_urls
|
|
||||||
if u.startswith(CIVITAI_DOWNLOAD_URL_PREFIXES)
|
|
||||||
]
|
|
||||||
non_civitai_urls = [
|
|
||||||
u
|
|
||||||
for u in download_urls
|
|
||||||
if not u.startswith(CIVITAI_DOWNLOAD_URL_PREFIXES)
|
|
||||||
]
|
|
||||||
download_urls = non_civitai_urls + civitai_urls
|
|
||||||
else:
|
|
||||||
download_url = file_info.get("downloadUrl")
|
|
||||||
if download_url:
|
|
||||||
download_urls.append(
|
|
||||||
normalize_civitai_download_url(download_url)
|
|
||||||
)
|
|
||||||
|
|
||||||
if not download_urls:
|
if not download_urls:
|
||||||
return {"success": False, "error": "No mirror URL found"}
|
return {"success": False, "error": "No mirror URL found"}
|
||||||
@@ -1803,6 +1837,9 @@ class DownloadManager:
|
|||||||
model_tags, model_type
|
model_tags, model_type
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if not first_tag:
|
||||||
|
first_tag = "no tags" # Default if no tags available
|
||||||
|
|
||||||
# Format the template with available data
|
# Format the template with available data
|
||||||
formatted_path = path_template
|
formatted_path = path_template
|
||||||
formatted_path = formatted_path.replace("{base_model}", mapped_base_model)
|
formatted_path = formatted_path.replace("{base_model}", mapped_base_model)
|
||||||
@@ -1818,6 +1855,15 @@ class DownloadManager:
|
|||||||
if model_type == "embedding":
|
if model_type == "embedding":
|
||||||
formatted_path = formatted_path.replace(" ", "_")
|
formatted_path = formatted_path.replace(" ", "_")
|
||||||
|
|
||||||
|
# Sanitize the resolved path to prevent path traversal:
|
||||||
|
# - Strip leading slashes (prevents os.path.join from treating path as absolute)
|
||||||
|
# - Collapse double slashes from empty placeholder substitutions
|
||||||
|
# - Strip trailing slashes for cleanliness
|
||||||
|
formatted_path = formatted_path.lstrip("/")
|
||||||
|
while "//" in formatted_path:
|
||||||
|
formatted_path = formatted_path.replace("//", "/")
|
||||||
|
formatted_path = formatted_path.rstrip("/")
|
||||||
|
|
||||||
return formatted_path
|
return formatted_path
|
||||||
|
|
||||||
async def _execute_download(
|
async def _execute_download(
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ class DownloadQueueService:
|
|||||||
_instance: Optional[DownloadQueueService] = None
|
_instance: Optional[DownloadQueueService] = None
|
||||||
_class_lock: asyncio.Lock = asyncio.Lock()
|
_class_lock: asyncio.Lock = asyncio.Lock()
|
||||||
|
|
||||||
_SCHEMA = """
|
_SCHEMA_TABLES = """
|
||||||
CREATE TABLE IF NOT EXISTS download_queue (
|
CREATE TABLE IF NOT EXISTS download_queue (
|
||||||
download_id TEXT PRIMARY KEY,
|
download_id TEXT PRIMARY KEY,
|
||||||
model_id INTEGER,
|
model_id INTEGER,
|
||||||
@@ -76,6 +76,11 @@ class DownloadQueueService:
|
|||||||
CREATE INDEX IF NOT EXISTS idx_dh_status ON download_history(status);
|
CREATE INDEX IF NOT EXISTS idx_dh_status ON download_history(status);
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
_CREATE_UNIQUE_INDEX = """
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS idx_dh_download_id
|
||||||
|
ON download_history(download_id) WHERE download_id IS NOT NULL;
|
||||||
|
"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def get_instance(cls) -> DownloadQueueService:
|
async def get_instance(cls) -> DownloadQueueService:
|
||||||
"""Return the singleton instance, creating it if necessary."""
|
"""Return the singleton instance, creating it if necessary."""
|
||||||
@@ -113,10 +118,39 @@ class DownloadQueueService:
|
|||||||
if self._schema_initialized:
|
if self._schema_initialized:
|
||||||
return
|
return
|
||||||
with self._connect() as conn:
|
with self._connect() as conn:
|
||||||
conn.executescript(self._SCHEMA)
|
conn.executescript(self._SCHEMA_TABLES)
|
||||||
|
|
||||||
|
# Creating the unique index on download_history.download_id can
|
||||||
|
# fail if pre-existing rows have duplicate values (e.g. from a
|
||||||
|
# previous version that lacked the index). Deduplicate first so
|
||||||
|
# that the migration does not crash on startup.
|
||||||
|
if not self._index_exists(conn, "idx_dh_download_id"):
|
||||||
|
self._remove_duplicate_download_ids(conn)
|
||||||
|
conn.executescript(self._CREATE_UNIQUE_INDEX)
|
||||||
|
|
||||||
conn.commit()
|
conn.commit()
|
||||||
self._schema_initialized = True
|
self._schema_initialized = True
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _index_exists(conn: sqlite3.Connection, name: str) -> bool:
|
||||||
|
return conn.execute(
|
||||||
|
"SELECT 1 FROM sqlite_master WHERE type='index' AND name=?",
|
||||||
|
(name,),
|
||||||
|
).fetchone() is not None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _remove_duplicate_download_ids(conn: sqlite3.Connection) -> None:
|
||||||
|
conn.execute("""
|
||||||
|
DELETE FROM download_history
|
||||||
|
WHERE id NOT IN (
|
||||||
|
SELECT MIN(id)
|
||||||
|
FROM download_history
|
||||||
|
WHERE download_id IS NOT NULL
|
||||||
|
GROUP BY download_id
|
||||||
|
)
|
||||||
|
AND download_id IS NOT NULL
|
||||||
|
""")
|
||||||
|
|
||||||
def get_database_path(self) -> str:
|
def get_database_path(self) -> str:
|
||||||
"""Return the resolved database file path."""
|
"""Return the resolved database file path."""
|
||||||
return self._db_path
|
return self._db_path
|
||||||
@@ -154,13 +188,23 @@ class DownloadQueueService:
|
|||||||
"""Insert a new download into the queue.
|
"""Insert a new download into the queue.
|
||||||
|
|
||||||
Returns the inserted row as a dict (or an empty dict if the
|
Returns the inserted row as a dict (or an empty dict if the
|
||||||
download_id already exists).
|
download_id already exists in the queue or has a terminal
|
||||||
|
record in history).
|
||||||
"""
|
"""
|
||||||
now = time.time()
|
now = time.time()
|
||||||
file_params_json = json.dumps(file_params) if file_params is not None else None
|
file_params_json = json.dumps(file_params) if file_params is not None else None
|
||||||
|
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
conn = self._get_conn()
|
conn = self._get_conn()
|
||||||
|
|
||||||
|
# Reject download_ids that already have a terminal record in history.
|
||||||
|
history_row = conn.execute(
|
||||||
|
"SELECT 1 FROM download_history WHERE download_id = ? LIMIT 1",
|
||||||
|
(download_id,),
|
||||||
|
).fetchone()
|
||||||
|
if history_row is not None:
|
||||||
|
return {}
|
||||||
|
|
||||||
conn.execute(
|
conn.execute(
|
||||||
"""
|
"""
|
||||||
INSERT OR IGNORE INTO download_queue (
|
INSERT OR IGNORE INTO download_queue (
|
||||||
@@ -380,7 +424,7 @@ class DownloadQueueService:
|
|||||||
)
|
)
|
||||||
conn.execute(
|
conn.execute(
|
||||||
"""
|
"""
|
||||||
INSERT INTO download_history (
|
INSERT OR IGNORE INTO download_history (
|
||||||
download_id, model_id, model_version_id, model_name,
|
download_id, model_id, model_version_id, model_name,
|
||||||
version_name, thumbnail_url, status, error, file_path,
|
version_name, thumbnail_url, status, error, file_path,
|
||||||
bytes_downloaded, total_bytes, completed_at
|
bytes_downloaded, total_bytes, completed_at
|
||||||
@@ -537,17 +581,27 @@ class DownloadQueueService:
|
|||||||
"offset": offset,
|
"offset": offset,
|
||||||
}
|
}
|
||||||
|
|
||||||
async def delete_history_item(self, id: int) -> bool:
|
async def delete_history_item(
|
||||||
"""Delete a single history entry by its *id*.
|
self, id: Optional[int] = None, download_id: Optional[str] = None
|
||||||
|
) -> bool:
|
||||||
|
"""Delete a single history entry by *download_id* (preferred) or *id*.
|
||||||
|
|
||||||
Returns ``True`` if a row was deleted.
|
Returns ``True`` if a row was deleted.
|
||||||
"""
|
"""
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
conn = self._get_conn()
|
conn = self._get_conn()
|
||||||
cursor = conn.execute(
|
if download_id:
|
||||||
"DELETE FROM download_history WHERE id = ?",
|
cursor = conn.execute(
|
||||||
(id,),
|
"DELETE FROM download_history WHERE download_id = ?",
|
||||||
)
|
(download_id,),
|
||||||
|
)
|
||||||
|
elif id is not None:
|
||||||
|
cursor = conn.execute(
|
||||||
|
"DELETE FROM download_history WHERE id = ?",
|
||||||
|
(id,),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
return False
|
||||||
conn.commit()
|
conn.commit()
|
||||||
return cursor.rowcount > 0
|
return cursor.rowcount > 0
|
||||||
|
|
||||||
@@ -604,21 +658,34 @@ class DownloadQueueService:
|
|||||||
# Retry
|
# Retry
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
async def retry_from_history(self, item_id: int) -> Optional[dict[str, Any]]:
|
async def retry_from_history(
|
||||||
|
self,
|
||||||
|
item_id: Optional[int] = None,
|
||||||
|
download_id: Optional[str] = None,
|
||||||
|
) -> Optional[dict[str, Any]]:
|
||||||
"""Re-queue a failed or canceled download from history.
|
"""Re-queue a failed or canceled download from history.
|
||||||
|
|
||||||
Looks up the history record by its primary key. If the status is
|
Looks up the history record by *download_id* (preferred) or
|
||||||
``failed`` or ``canceled`` a new queue entry is created with the
|
*item_id*. If the status is ``failed`` or ``canceled`` a new
|
||||||
same model metadata and a fresh download id, and the original
|
queue entry is created with the same model metadata and a fresh
|
||||||
history entry is **deleted** to prevent exponential growth when
|
download id, and the original history entry is **deleted** to
|
||||||
the retried item is later canceled or fails again and re-retried.
|
prevent exponential growth when the retried item is later
|
||||||
|
canceled or fails again and re-retried.
|
||||||
"""
|
"""
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
conn = self._get_conn()
|
conn = self._get_conn()
|
||||||
row = conn.execute(
|
if download_id:
|
||||||
"SELECT * FROM download_history WHERE id = ?",
|
row = conn.execute(
|
||||||
(item_id,),
|
"SELECT * FROM download_history WHERE download_id = ?",
|
||||||
).fetchone()
|
(download_id,),
|
||||||
|
).fetchone()
|
||||||
|
elif item_id is not None:
|
||||||
|
row = conn.execute(
|
||||||
|
"SELECT * FROM download_history WHERE id = ?",
|
||||||
|
(item_id,),
|
||||||
|
).fetchone()
|
||||||
|
else:
|
||||||
|
return None
|
||||||
if row is None:
|
if row is None:
|
||||||
return None
|
return None
|
||||||
status = str(row["status"])
|
status = str(row["status"])
|
||||||
@@ -650,7 +717,7 @@ class DownloadQueueService:
|
|||||||
)
|
)
|
||||||
conn.execute(
|
conn.execute(
|
||||||
"DELETE FROM download_history WHERE id = ?",
|
"DELETE FROM download_history WHERE id = ?",
|
||||||
(item_id,),
|
(row["id"],),
|
||||||
)
|
)
|
||||||
conn.commit()
|
conn.commit()
|
||||||
queued = conn.execute(
|
queued = conn.execute(
|
||||||
|
|||||||
+19
-10
@@ -270,14 +270,14 @@ class Downloader:
|
|||||||
|
|
||||||
Note: This is private and caller MUST hold self._session_lock.
|
Note: This is private and caller MUST hold self._session_lock.
|
||||||
"""
|
"""
|
||||||
# Close existing session if any
|
# Snapshot and clear old session reference before creating the new
|
||||||
if self._session is not None:
|
# one. This ensures self._session is always valid (or None, which
|
||||||
try:
|
# triggers a fresh creation) and avoids a race where concurrent
|
||||||
await self._session.close()
|
# requests hold a reference to a session whose connector has been
|
||||||
except Exception as e: # pragma: no cover
|
# torn down by a premature close() call — the root cause of the
|
||||||
logger.warning(f"Error closing previous session: {e}")
|
# intermittent "NoneType has no attribute connect" crash.
|
||||||
finally:
|
old_session = self._session
|
||||||
self._session = None
|
self._session = None
|
||||||
|
|
||||||
# Check for app-level proxy settings
|
# Check for app-level proxy settings
|
||||||
proxy_url = None # http(s) proxy, passed via the per-request `proxy=` kwarg
|
proxy_url = None # http(s) proxy, passed via the per-request `proxy=` kwarg
|
||||||
@@ -372,6 +372,13 @@ class Downloader:
|
|||||||
self._proxy_url = proxy_url
|
self._proxy_url = proxy_url
|
||||||
self._session_created_at = datetime.now()
|
self._session_created_at = datetime.now()
|
||||||
|
|
||||||
|
# Close the previous session now that the replacement is live.
|
||||||
|
if old_session is not None:
|
||||||
|
try:
|
||||||
|
await old_session.close()
|
||||||
|
except Exception as e: # pragma: no cover
|
||||||
|
logger.warning(f"Error closing previous session: {e}")
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Created new HTTP session with proxy settings. App-level proxy: %s, System-level proxy (trust_env): %s",
|
"Created new HTTP session with proxy settings. App-level proxy: %s, System-level proxy (trust_env): %s",
|
||||||
bool(proxy_url),
|
bool(proxy_url),
|
||||||
@@ -753,7 +760,8 @@ class Downloader:
|
|||||||
else:
|
else:
|
||||||
resume_offset = 0
|
resume_offset = 0
|
||||||
total_size = 0
|
total_size = 0
|
||||||
await self._create_session()
|
async with self._session_lock:
|
||||||
|
await self._create_session()
|
||||||
continue
|
continue
|
||||||
|
|
||||||
return False, integrity_error
|
return False, integrity_error
|
||||||
@@ -843,7 +851,8 @@ class Downloader:
|
|||||||
logger.info(f"Will resume from byte {resume_offset}")
|
logger.info(f"Will resume from byte {resume_offset}")
|
||||||
|
|
||||||
# Refresh session to get new connection
|
# Refresh session to get new connection
|
||||||
await self._create_session()
|
async with self._session_lock:
|
||||||
|
await self._create_session()
|
||||||
continue
|
continue
|
||||||
else:
|
else:
|
||||||
logger.error(f"Max retries exceeded for download: {e}")
|
logger.error(f"Max retries exceeded for download: {e}")
|
||||||
|
|||||||
@@ -25,3 +25,21 @@ class ResourceNotFoundError(RuntimeError):
|
|||||||
|
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class LLMNotConfiguredError(RuntimeError):
|
||||||
|
"""Raised when an LLM-dependent operation is attempted but no provider is configured."""
|
||||||
|
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class LLMRateLimitError(RateLimitError):
|
||||||
|
"""Raised when the LLM provider rejects a request due to rate limiting."""
|
||||||
|
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class LLMResponseError(RuntimeError):
|
||||||
|
"""Raised when the LLM returns an unparseable or schema-invalid response."""
|
||||||
|
|
||||||
|
pass
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,729 @@
|
|||||||
|
"""Centralized LLM API client with BYOK (bring-your-own-key) provider support.
|
||||||
|
|
||||||
|
Reads provider configuration from :class:`SettingsManager` and makes
|
||||||
|
OpenAI-compatible ``/chat/completions`` calls. Supports any provider that
|
||||||
|
implements the OpenAI Chat Completions API surface area (OpenAI, Ollama,
|
||||||
|
vLLM, LM Studio, etc.).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
import aiohttp
|
||||||
|
|
||||||
|
from .errors import LLMNotConfiguredError, LLMRateLimitError, LLMResponseError
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Model catalog sourced from opencode's maintained model registry.
|
||||||
|
# maps provider_id -> list of model IDs.
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
_MODEL_CATALOG_URL = "https://models.dev/api.json"
|
||||||
|
|
||||||
|
# In-memory cache: maps provider slug -> list of model ID strings.
|
||||||
|
_catalog_cache: Optional[Dict[str, List[str]]] = None
|
||||||
|
|
||||||
|
# Per-model max output token limits parsed from the catalog.
|
||||||
|
# ``{provider_id: {model_id: max_output_tokens}}``.
|
||||||
|
_model_output_limits: Dict[str, Dict[str, int]] = {}
|
||||||
|
|
||||||
|
_CATALOG_TIMEOUT = aiohttp.ClientTimeout(total=30)
|
||||||
|
|
||||||
|
|
||||||
|
async def _load_model_catalog() -> Dict[str, List[str]]:
|
||||||
|
"""Fetch and parse the model catalog.
|
||||||
|
|
||||||
|
Returns ``{provider_id: [model_id, ...]}`` and also populates
|
||||||
|
:data:`_model_output_limits` with per-model ``limit.output`` values
|
||||||
|
for use by :func:`_get_model_max_output`.
|
||||||
|
|
||||||
|
The JSON at ``_MODEL_CATALOG_URL`` is a dict keyed by provider slug; each
|
||||||
|
value has a ``models`` sub-dict keyed by model ID. The result is cached
|
||||||
|
in memory after the first successful fetch.
|
||||||
|
Subsequent calls return the cached data immediately.
|
||||||
|
"""
|
||||||
|
global _catalog_cache, _model_output_limits
|
||||||
|
if _catalog_cache is not None:
|
||||||
|
return _catalog_cache
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with aiohttp.ClientSession(timeout=_CATALOG_TIMEOUT) as session:
|
||||||
|
async with session.get(_MODEL_CATALOG_URL) as resp:
|
||||||
|
if resp.status != 200:
|
||||||
|
logger.warning("Model catalog returned HTTP %s", resp.status)
|
||||||
|
return _catalog_cache or {}
|
||||||
|
data = await resp.json()
|
||||||
|
except (aiohttp.ClientError, asyncio.TimeoutError, json.JSONDecodeError) as exc:
|
||||||
|
logger.warning("Failed to fetch model catalog: %s", exc)
|
||||||
|
return _catalog_cache or {}
|
||||||
|
|
||||||
|
if not isinstance(data, dict):
|
||||||
|
logger.warning("Model catalog is not a dict, got %s", type(data).__name__)
|
||||||
|
return _catalog_cache or {}
|
||||||
|
|
||||||
|
result: Dict[str, List[str]] = {}
|
||||||
|
output_limits: Dict[str, Dict[str, int]] = {}
|
||||||
|
for provider_id, provider_info in data.items():
|
||||||
|
if not isinstance(provider_info, dict):
|
||||||
|
continue
|
||||||
|
models_dict = provider_info.get("models")
|
||||||
|
if not isinstance(models_dict, dict):
|
||||||
|
continue
|
||||||
|
model_ids: List[str] = []
|
||||||
|
provider_limits: Dict[str, int] = {}
|
||||||
|
for mid, model_info in models_dict.items():
|
||||||
|
if not isinstance(mid, str):
|
||||||
|
continue
|
||||||
|
model_ids.append(mid)
|
||||||
|
if isinstance(model_info, dict):
|
||||||
|
limit = model_info.get("limit")
|
||||||
|
if isinstance(limit, dict):
|
||||||
|
output = limit.get("output")
|
||||||
|
if isinstance(output, (int, float)) and output > 0:
|
||||||
|
provider_limits[mid] = int(output)
|
||||||
|
if model_ids:
|
||||||
|
result[provider_id] = model_ids
|
||||||
|
if provider_limits:
|
||||||
|
output_limits[provider_id] = provider_limits
|
||||||
|
|
||||||
|
_catalog_cache = result
|
||||||
|
_model_output_limits = output_limits
|
||||||
|
logger.debug(
|
||||||
|
"Loaded model catalog: %d providers, %d total models "
|
||||||
|
"(%d providers have output limits)",
|
||||||
|
len(result),
|
||||||
|
sum(len(m) for m in result.values()),
|
||||||
|
len(output_limits),
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _get_model_max_output(provider: str, model: str) -> Optional[int]:
|
||||||
|
"""Return the model's max output token limit from the catalog, or ``None``.
|
||||||
|
|
||||||
|
Returns ``None`` when the provider or model is not found in the catalog
|
||||||
|
(e.g. local Ollama models, custom models, or user-typed model names).
|
||||||
|
Callers should fall back to a safe default.
|
||||||
|
"""
|
||||||
|
return _model_output_limits.get(provider, {}).get(model)
|
||||||
|
|
||||||
|
|
||||||
|
# Short timeout for Ollama's local API
|
||||||
|
_OLLAMA_API_TIMEOUT = aiohttp.ClientTimeout(total=8)
|
||||||
|
|
||||||
|
|
||||||
|
async def fetch_ollama_models(api_base: str) -> List[str]:
|
||||||
|
"""Fetch locally available models from a running Ollama instance.
|
||||||
|
|
||||||
|
Uses Ollama's OpenAI-compatible ``GET {api_base}/models`` endpoint.
|
||||||
|
Returns an empty list if Ollama is not reachable (not running).
|
||||||
|
"""
|
||||||
|
url = f"{api_base.rstrip('/')}/models"
|
||||||
|
try:
|
||||||
|
async with aiohttp.ClientSession(timeout=_OLLAMA_API_TIMEOUT) as session:
|
||||||
|
async with session.get(url) as resp:
|
||||||
|
if resp.status != 200:
|
||||||
|
logger.debug("Ollama API returned HTTP %s from %s", resp.status, api_base)
|
||||||
|
return []
|
||||||
|
data = await resp.json()
|
||||||
|
except (aiohttp.ClientError, asyncio.TimeoutError, json.JSONDecodeError) as exc:
|
||||||
|
logger.debug("Ollama not reachable at %s: %s", api_base, exc)
|
||||||
|
return []
|
||||||
|
|
||||||
|
raw = data.get("data") if isinstance(data, dict) else None
|
||||||
|
if not isinstance(raw, list):
|
||||||
|
return []
|
||||||
|
|
||||||
|
return [
|
||||||
|
str(entry["id"]) for entry in raw
|
||||||
|
if isinstance(entry, dict) and isinstance(entry.get("id"), str)
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
async def get_provider_model_ids(provider_id: str) -> List[str]:
|
||||||
|
"""Return the list of known model IDs for *provider_id* from the catalog.
|
||||||
|
|
||||||
|
The catalog is loaded on first call and cached thereafter. If the
|
||||||
|
provider is not found an empty list is returned (never raises).
|
||||||
|
"""
|
||||||
|
catalog = await _load_model_catalog()
|
||||||
|
return catalog.get(provider_id, [])
|
||||||
|
|
||||||
|
|
||||||
|
async def get_all_provider_models(
|
||||||
|
provider_ids: List[str],
|
||||||
|
) -> Dict[str, List[str]]:
|
||||||
|
"""Return model lists for a subset of providers in one call.
|
||||||
|
|
||||||
|
Loads the catalog (cached) and returns only the requested providers.
|
||||||
|
Handy for embedding lightweight data into the template context.
|
||||||
|
"""
|
||||||
|
catalog = await _load_model_catalog()
|
||||||
|
return {
|
||||||
|
pid: catalog.get(pid, [])
|
||||||
|
for pid in provider_ids
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# Provider preset definitions.
|
||||||
|
# Each entry contains display metadata and defaults for the UI.
|
||||||
|
# The key is the internal provider id stored in ``llm_provider``.
|
||||||
|
# Models are NOT listed here — they come from the opencode model catalog at
|
||||||
|
# runtime (see :func:`get_provider_model_ids`).
|
||||||
|
PROVIDER_PRESETS: Dict[str, Dict[str, Any]] = {
|
||||||
|
"openai": {
|
||||||
|
"name": "OpenAI",
|
||||||
|
"api_base": "https://api.openai.com/v1",
|
||||||
|
"requires_key": True,
|
||||||
|
},
|
||||||
|
"ollama": {
|
||||||
|
"name": "Ollama (local)",
|
||||||
|
"api_base": "http://localhost:11434/v1",
|
||||||
|
"requires_key": False,
|
||||||
|
},
|
||||||
|
"deepseek": {
|
||||||
|
"name": "DeepSeek",
|
||||||
|
"api_base": "https://api.deepseek.com/v1",
|
||||||
|
"requires_key": True,
|
||||||
|
},
|
||||||
|
"groq": {
|
||||||
|
"name": "Groq",
|
||||||
|
"api_base": "https://api.groq.com/openai/v1",
|
||||||
|
"requires_key": True,
|
||||||
|
},
|
||||||
|
"openrouter": {
|
||||||
|
"name": "OpenRouter",
|
||||||
|
"api_base": "https://openrouter.ai/api/v1",
|
||||||
|
"requires_key": True,
|
||||||
|
},
|
||||||
|
"opencode-go": {
|
||||||
|
"name": "OpenCode Go",
|
||||||
|
"api_base": "https://opencode.ai/zen/go/v1",
|
||||||
|
"requires_key": True,
|
||||||
|
},
|
||||||
|
# "custom" is handled specially (no preset api_base, requires user input)
|
||||||
|
}
|
||||||
|
|
||||||
|
# Legacy lookup derived from PROVIDER_PRESETS for backward compat.
|
||||||
|
_PROVIDER_DEFAULTS: Dict[str, str] = {
|
||||||
|
pid: info["api_base"]
|
||||||
|
for pid, info in PROVIDER_PRESETS.items()
|
||||||
|
if info.get("api_base")
|
||||||
|
}
|
||||||
|
|
||||||
|
# Request timeout for LLM calls (seconds)
|
||||||
|
_LLM_TIMEOUT = aiohttp.ClientTimeout(total=120)
|
||||||
|
|
||||||
|
|
||||||
|
class LLMService:
|
||||||
|
"""Centralized LLM API client.
|
||||||
|
|
||||||
|
All LLM-based enrichment features call through this service so
|
||||||
|
that BYOK config, retry logic, and error handling live in one place.
|
||||||
|
"""
|
||||||
|
|
||||||
|
_instance: Optional["LLMService"] = None
|
||||||
|
_lock: asyncio.Lock = asyncio.Lock()
|
||||||
|
|
||||||
|
def __init__(self, settings_service) -> None:
|
||||||
|
self._settings = settings_service
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Singleton access
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def get_instance(cls) -> "LLMService":
|
||||||
|
"""Return the lazily-initialised global ``LLMService`` instance."""
|
||||||
|
|
||||||
|
if cls._instance is None:
|
||||||
|
async with cls._lock:
|
||||||
|
if cls._instance is None:
|
||||||
|
from .settings_manager import get_settings_manager
|
||||||
|
|
||||||
|
cls._instance = cls(get_settings_manager())
|
||||||
|
# Start preloading the model catalog in the background so
|
||||||
|
# the settings UI never blocks on it. The catalog is
|
||||||
|
# cached after the first fetch (see _load_model_catalog).
|
||||||
|
asyncio.create_task(_load_model_catalog())
|
||||||
|
return cls._instance
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def reset_instance(cls) -> None:
|
||||||
|
"""Reset the cached singleton — primarily for tests."""
|
||||||
|
|
||||||
|
cls._instance = None
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Configuration helpers
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _get_config(self) -> Dict[str, Any]:
|
||||||
|
"""Read the current LLM configuration from settings."""
|
||||||
|
|
||||||
|
return {
|
||||||
|
"provider": self._settings.get("llm_provider", "openai"),
|
||||||
|
"api_key": self._settings.get("llm_api_key", ""),
|
||||||
|
"api_base": self._settings.get("llm_api_base", ""),
|
||||||
|
"model": self._settings.get("llm_model", ""),
|
||||||
|
}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _provider_requires_key(provider: str) -> bool:
|
||||||
|
"""Return ``False`` when the given provider id does not need an API key."""
|
||||||
|
preset = PROVIDER_PRESETS.get(provider, {})
|
||||||
|
return bool(preset.get("requires_key", True))
|
||||||
|
|
||||||
|
def is_configured(self) -> bool:
|
||||||
|
"""Return ``True`` when the LLM provider is minimally configured.
|
||||||
|
|
||||||
|
A provider is considered configured when ``llm_model`` is set,
|
||||||
|
an API key is configured for providers that require one (e.g.
|
||||||
|
Ollama does not), and an API base URL is set for providers that
|
||||||
|
have no preset default (e.g. ``custom``).
|
||||||
|
"""
|
||||||
|
|
||||||
|
cfg = self._get_config()
|
||||||
|
has_model = bool(cfg["model"])
|
||||||
|
has_key = bool(cfg["api_key"]) or not self._provider_requires_key(cfg["provider"])
|
||||||
|
has_base = bool(cfg["api_base"]) or bool(_PROVIDER_DEFAULTS.get(cfg["provider"]))
|
||||||
|
return has_model and has_key and has_base
|
||||||
|
|
||||||
|
def _resolve_api_base(self, provider: str, api_base: str) -> str:
|
||||||
|
"""Resolve the API base URL for the given provider.
|
||||||
|
|
||||||
|
If ``api_base`` is explicitly set (non-empty), it takes priority.
|
||||||
|
Otherwise the default from :data:`PROVIDER_PRESETS` is used.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if api_base:
|
||||||
|
return api_base.rstrip("/")
|
||||||
|
return _PROVIDER_DEFAULTS.get(provider, "").rstrip("/")
|
||||||
|
|
||||||
|
def _build_headers(self, api_key: str) -> Dict[str, str]:
|
||||||
|
"""Build HTTP headers for the LLM API request."""
|
||||||
|
|
||||||
|
headers = {"Content-Type": "application/json"}
|
||||||
|
if api_key:
|
||||||
|
headers["Authorization"] = f"Bearer {api_key}"
|
||||||
|
return headers
|
||||||
|
|
||||||
|
def _ensure_configured(self) -> Dict[str, Any]:
|
||||||
|
"""Validate configuration and return it, or raise.
|
||||||
|
|
||||||
|
A provider is considered configured when ``llm_model`` is set,
|
||||||
|
an API key is configured for providers that require one, and
|
||||||
|
an API base URL is set for providers without a preset default.
|
||||||
|
"""
|
||||||
|
|
||||||
|
cfg = self._get_config()
|
||||||
|
has_model = bool(cfg["model"])
|
||||||
|
needs_key = self._provider_requires_key(cfg["provider"])
|
||||||
|
has_key = bool(cfg["api_key"]) or not needs_key
|
||||||
|
has_base = bool(cfg["api_base"]) or bool(_PROVIDER_DEFAULTS.get(cfg["provider"]))
|
||||||
|
if not (has_model and has_key and has_base):
|
||||||
|
parts = []
|
||||||
|
if not has_model:
|
||||||
|
parts.append("No LLM model specified")
|
||||||
|
if not has_key and needs_key:
|
||||||
|
parts.append("No LLM API key configured")
|
||||||
|
if not has_base:
|
||||||
|
parts.append(
|
||||||
|
f"No API base URL for provider '{cfg['provider']}'"
|
||||||
|
)
|
||||||
|
detail = "; ".join(parts) if parts else "LLM provider is not configured"
|
||||||
|
raise LLMNotConfiguredError(
|
||||||
|
f"{detail}. Configure it in Settings → AI Provider."
|
||||||
|
)
|
||||||
|
return cfg
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Core API call
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def chat_completion(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
messages: List[Dict[str, str]],
|
||||||
|
model: Optional[str] = None,
|
||||||
|
temperature: float = 0.3,
|
||||||
|
response_format: Optional[Dict[str, Any]] = None,
|
||||||
|
max_tokens: Optional[int] = None,
|
||||||
|
retry_on_rate_limit: bool = True,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""Call the configured LLM provider's ``/chat/completions`` endpoint.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
messages: OpenAI-format message list
|
||||||
|
model: Override the configured model name
|
||||||
|
temperature: Sampling temperature
|
||||||
|
response_format: Optional ``{"type": "json_object"}`` for structured output
|
||||||
|
max_tokens: Optional max output tokens
|
||||||
|
retry_on_rate_limit: Retry once after a 429 with backoff
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict with ``content`` (str), ``usage`` (dict), ``model`` (str)
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
LLMNotConfiguredError: Provider not enabled / missing config
|
||||||
|
LLMRateLimitError: Rate limited and retry exhausted
|
||||||
|
LLMResponseError: Non-200 response or parse failure
|
||||||
|
"""
|
||||||
|
|
||||||
|
cfg = self._ensure_configured()
|
||||||
|
api_base = self._resolve_api_base(cfg["provider"], cfg["api_base"])
|
||||||
|
model_name = model or cfg["model"]
|
||||||
|
|
||||||
|
is_ollama = cfg["provider"] == "ollama"
|
||||||
|
|
||||||
|
if is_ollama:
|
||||||
|
# Use Ollama's native /api/chat endpoint which does NOT expose
|
||||||
|
# a separate reasoning/thinking field (the model's full output
|
||||||
|
# lands directly in message.content). The OpenAI-compatible
|
||||||
|
# endpoint splits thinking into the "reasoning" field, making
|
||||||
|
# content empty when thinking consumes all available tokens.
|
||||||
|
base = api_base.rstrip("/")
|
||||||
|
if base.endswith("/v1"):
|
||||||
|
base = base[:-3]
|
||||||
|
url = f"{base}/api/chat"
|
||||||
|
else:
|
||||||
|
url = f"{api_base}/chat/completions"
|
||||||
|
|
||||||
|
payload: Dict[str, Any]
|
||||||
|
if is_ollama:
|
||||||
|
payload = {
|
||||||
|
"model": model_name,
|
||||||
|
"messages": messages,
|
||||||
|
"stream": False,
|
||||||
|
# Suppress separate thinking trace — thinking still happens
|
||||||
|
# internally (accuracy preserved) but output goes directly to
|
||||||
|
# message.content instead of being split across content +
|
||||||
|
# thinking. Without this the model can exhaust num_predict
|
||||||
|
# on thinking alone and leave content empty.
|
||||||
|
"think": False,
|
||||||
|
"options": {
|
||||||
|
"temperature": temperature,
|
||||||
|
# 8K context is sufficient for metadata enrichment
|
||||||
|
# (prompt ~2-5K, output ~0.2-1K tokens). The old 32K
|
||||||
|
# value was excessive for this use case and increased
|
||||||
|
# Ollama VRAM usage unnecessarily.
|
||||||
|
"num_ctx": 8192,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if response_format is not None:
|
||||||
|
payload["format"] = "json"
|
||||||
|
if max_tokens is not None:
|
||||||
|
payload["options"]["num_predict"] = max_tokens
|
||||||
|
else:
|
||||||
|
payload = {
|
||||||
|
"model": model_name,
|
||||||
|
"messages": messages,
|
||||||
|
"temperature": temperature,
|
||||||
|
}
|
||||||
|
if response_format is not None:
|
||||||
|
payload["response_format"] = response_format
|
||||||
|
if max_tokens is not None:
|
||||||
|
payload["max_tokens"] = max_tokens
|
||||||
|
|
||||||
|
if is_ollama:
|
||||||
|
logger.info(
|
||||||
|
"Ollama request: model=%s num_ctx=%s num_predict=%s format=%s think=%s",
|
||||||
|
payload.get("model"),
|
||||||
|
payload.get("options", {}).get("num_ctx"),
|
||||||
|
payload.get("options", {}).get("num_predict"),
|
||||||
|
payload.get("format", "none"),
|
||||||
|
payload.get("think"),
|
||||||
|
)
|
||||||
|
|
||||||
|
headers = self._build_headers(cfg["api_key"])
|
||||||
|
|
||||||
|
attempt = 0
|
||||||
|
max_attempts = 2 if retry_on_rate_limit else 1
|
||||||
|
while attempt < max_attempts:
|
||||||
|
attempt += 1
|
||||||
|
try:
|
||||||
|
async with aiohttp.ClientSession(timeout=_LLM_TIMEOUT) as session:
|
||||||
|
async with session.post(
|
||||||
|
url, json=payload, headers=headers
|
||||||
|
) as resp:
|
||||||
|
if resp.status == 429:
|
||||||
|
if attempt < max_attempts:
|
||||||
|
retry_after = float(
|
||||||
|
resp.headers.get("Retry-After", "5")
|
||||||
|
)
|
||||||
|
logger.warning(
|
||||||
|
"LLM rate limited, retrying after %.1fs",
|
||||||
|
retry_after,
|
||||||
|
)
|
||||||
|
await asyncio.sleep(retry_after)
|
||||||
|
continue
|
||||||
|
raise LLMRateLimitError(
|
||||||
|
f"LLM provider rate limited (HTTP 429)",
|
||||||
|
provider=cfg["provider"],
|
||||||
|
)
|
||||||
|
|
||||||
|
if resp.status != 200:
|
||||||
|
body = await resp.text()
|
||||||
|
raise LLMResponseError(
|
||||||
|
f"LLM API returned HTTP {resp.status}: "
|
||||||
|
f"{body[:500]}"
|
||||||
|
)
|
||||||
|
|
||||||
|
data = await resp.json()
|
||||||
|
|
||||||
|
except aiohttp.ClientError as exc:
|
||||||
|
raise LLMResponseError(f"Network error calling LLM API: {exc}") from exc
|
||||||
|
|
||||||
|
# Parse response
|
||||||
|
try:
|
||||||
|
if is_ollama:
|
||||||
|
content = (data.get("message") or {}).get("content") or ""
|
||||||
|
usage = {"completion_tokens": data.get("eval_count", 0)}
|
||||||
|
finish_reason = data.get("done_reason", "")
|
||||||
|
if not content:
|
||||||
|
logger.warning(
|
||||||
|
"LLM returned empty content. Provider=ollama, "
|
||||||
|
"done_reason=%s, eval_count=%s",
|
||||||
|
finish_reason,
|
||||||
|
data.get("eval_count", 0),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
content = data["choices"][0]["message"].get("content") or ""
|
||||||
|
usage = data.get("usage", {})
|
||||||
|
if not content:
|
||||||
|
logger.warning(
|
||||||
|
"LLM returned empty content. Full response truncated: %s",
|
||||||
|
json.dumps(data, ensure_ascii=False)[:1000],
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"content": content,
|
||||||
|
"usage": usage,
|
||||||
|
"model": data.get("model", model_name),
|
||||||
|
}
|
||||||
|
except (KeyError, IndexError) as exc:
|
||||||
|
raise LLMResponseError(
|
||||||
|
f"Unexpected LLM response structure: {json.dumps(data)[:500]}"
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
# Should not reach here, but satisfy type checker
|
||||||
|
raise LLMRateLimitError("Rate limit retry exhausted", provider=cfg["provider"])
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Structured output convenience
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def chat_completion_json(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
system_prompt: str,
|
||||||
|
user_prompt: str,
|
||||||
|
model: Optional[str] = None,
|
||||||
|
temperature: float = 0.3,
|
||||||
|
max_tokens: Optional[int] = None,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""Call the LLM with ``response_format=json_object`` and return parsed JSON.
|
||||||
|
|
||||||
|
``max_tokens`` is resolved in this order:
|
||||||
|
1. Explicit caller-supplied ``max_tokens``
|
||||||
|
2. Per-model ``limit.output`` from the model catalog
|
||||||
|
3. A safe default of 4096 (sufficient for metadata enrichment)
|
||||||
|
|
||||||
|
If the response content is empty or not valid JSON, attempts
|
||||||
|
:func:`_try_salvage_json` before raising.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
system_prompt: System-level instructions
|
||||||
|
user_prompt: User-level query
|
||||||
|
model: Override the configured model name
|
||||||
|
temperature: Sampling temperature
|
||||||
|
max_tokens: Optional max output tokens
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Parsed JSON dict from the LLM response
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
LLMNotConfiguredError: Provider not configured
|
||||||
|
LLMRateLimitError: Rate limited
|
||||||
|
LLMResponseError: Empty response or JSON parse failure
|
||||||
|
"""
|
||||||
|
|
||||||
|
messages = [
|
||||||
|
{"role": "system", "content": system_prompt},
|
||||||
|
{"role": "user", "content": user_prompt},
|
||||||
|
]
|
||||||
|
|
||||||
|
# Resolve max_tokens: caller override → catalog lookup → safe default
|
||||||
|
if max_tokens is None:
|
||||||
|
cfg = self._get_config()
|
||||||
|
effective_max = _get_model_max_output(cfg["provider"], cfg["model"])
|
||||||
|
else:
|
||||||
|
effective_max = max_tokens
|
||||||
|
if effective_max is None:
|
||||||
|
effective_max = 4096
|
||||||
|
|
||||||
|
# Use json_schema (not json_object) for broader provider compatibility:
|
||||||
|
# LM Studio and some other OpenAI-compatible servers reject
|
||||||
|
# json_object but accept json_schema. {"type": "object"} is
|
||||||
|
# functionally equivalent — it accepts any JSON object without
|
||||||
|
# constraining specific fields.
|
||||||
|
response_format = {
|
||||||
|
"type": "json_schema",
|
||||||
|
"json_schema": {
|
||||||
|
"name": "metadata",
|
||||||
|
"schema": {"type": "object"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = await self.chat_completion(
|
||||||
|
messages=messages,
|
||||||
|
model=model,
|
||||||
|
temperature=temperature,
|
||||||
|
response_format=response_format,
|
||||||
|
max_tokens=effective_max,
|
||||||
|
)
|
||||||
|
except LLMResponseError as e:
|
||||||
|
# Only fall back when the provider rejects the response_format
|
||||||
|
# type value (e.g. "'response_format.type' must be..."). Avoid
|
||||||
|
# catching unrelated 400 errors whose body happens to mention
|
||||||
|
# "response_format" (e.g. "model does not support
|
||||||
|
# response_format restrictions on this endpoint").
|
||||||
|
if "'response_format.type'" not in str(e).lower():
|
||||||
|
raise
|
||||||
|
logger.info(
|
||||||
|
"Provider rejected response_format, retrying without it. "
|
||||||
|
"Falling back to prompt-only JSON mode. Error: %s",
|
||||||
|
e,
|
||||||
|
)
|
||||||
|
result = await self.chat_completion(
|
||||||
|
messages=messages,
|
||||||
|
model=model,
|
||||||
|
temperature=temperature,
|
||||||
|
response_format=None,
|
||||||
|
max_tokens=effective_max,
|
||||||
|
)
|
||||||
|
|
||||||
|
content = result.get("content", "") or ""
|
||||||
|
if not content:
|
||||||
|
raise LLMResponseError(
|
||||||
|
"LLM returned empty content. "
|
||||||
|
f"Raw response: {json.dumps(result)[:500]}"
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
parsed = json.loads(content)
|
||||||
|
logger.debug(
|
||||||
|
"LLM raw content: %s",
|
||||||
|
json.dumps(parsed, ensure_ascii=False)[:2000],
|
||||||
|
)
|
||||||
|
return parsed
|
||||||
|
except (json.JSONDecodeError, TypeError) as exc:
|
||||||
|
logger.info(
|
||||||
|
"LLM raw response (first 800 chars): %s",
|
||||||
|
content[:800],
|
||||||
|
)
|
||||||
|
|
||||||
|
# Last resort: attempt to salvage partial/truncated JSON
|
||||||
|
salvaged = _try_salvage_json(content)
|
||||||
|
if salvaged is not None:
|
||||||
|
logger.warning(
|
||||||
|
"LLM JSON salvaged from partial content (%d chars raw)",
|
||||||
|
len(content),
|
||||||
|
)
|
||||||
|
return salvaged
|
||||||
|
|
||||||
|
raise LLMResponseError(
|
||||||
|
f"LLM response could not be parsed as JSON: {content[:200]}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _try_salvage_json(raw: str) -> Dict[str, Any] | None:
|
||||||
|
"""Attempt to repair and parse a truncated JSON string.
|
||||||
|
|
||||||
|
Handles common truncation patterns:
|
||||||
|
|
||||||
|
* Incomplete string value at the end (``"foo`` → ``"foo"``)
|
||||||
|
* Missing closing ``}`` or ``]`` (respecting nesting order)
|
||||||
|
* Trailing comma before closing bracket
|
||||||
|
* Extra text after the JSON object (e.g. markdown fences)
|
||||||
|
|
||||||
|
Returns the parsed dict on success, ``None`` if repair is impossible.
|
||||||
|
"""
|
||||||
|
if not raw:
|
||||||
|
return None
|
||||||
|
|
||||||
|
text = raw.strip()
|
||||||
|
|
||||||
|
# Strip markdown fences if the LLM wrapped the JSON
|
||||||
|
if text.startswith("```"):
|
||||||
|
end = text.find("\n")
|
||||||
|
text = text[end + 1:] if end != -1 else text[3:]
|
||||||
|
if text.endswith("```"):
|
||||||
|
text = text[:-3].rstrip()
|
||||||
|
|
||||||
|
# Find the first '{' and strip everything before it
|
||||||
|
start = text.find("{")
|
||||||
|
if start == -1:
|
||||||
|
return None
|
||||||
|
text = text[start:]
|
||||||
|
|
||||||
|
# Try to close an incomplete string at the end (e.g. ``"https://huggingf``)
|
||||||
|
# Pattern: ends mid-string (last quote is open)
|
||||||
|
if text.count('"') % 2 == 1:
|
||||||
|
text += '"'
|
||||||
|
|
||||||
|
# Ensure trailing commas before closing braces work
|
||||||
|
text = _strip_trailing_commas(text)
|
||||||
|
|
||||||
|
# Walk through the text character by character to find unclosed
|
||||||
|
# brackets and close them in the correct (LIFO) order.
|
||||||
|
# We ignore brackets inside quoted strings.
|
||||||
|
stack: list[str] = []
|
||||||
|
in_string = False
|
||||||
|
escape = False
|
||||||
|
for ch in text:
|
||||||
|
if escape:
|
||||||
|
escape = False
|
||||||
|
continue
|
||||||
|
if ch == "\\":
|
||||||
|
escape = True
|
||||||
|
continue
|
||||||
|
if ch == '"':
|
||||||
|
in_string = not in_string
|
||||||
|
continue
|
||||||
|
if in_string:
|
||||||
|
continue
|
||||||
|
if ch in ("{", "["):
|
||||||
|
stack.append(ch)
|
||||||
|
elif ch == "}":
|
||||||
|
if stack and stack[-1] == "{":
|
||||||
|
stack.pop()
|
||||||
|
else:
|
||||||
|
return None # Unmatched closer — unrecoverable
|
||||||
|
elif ch == "]":
|
||||||
|
if stack and stack[-1] == "[":
|
||||||
|
stack.pop()
|
||||||
|
else:
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Close remaining open brackets in reverse order
|
||||||
|
for opener in reversed(stack):
|
||||||
|
text += "}" if opener == "{" else "]"
|
||||||
|
|
||||||
|
try:
|
||||||
|
return json.loads(text)
|
||||||
|
except (json.JSONDecodeError, ValueError):
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _strip_trailing_commas(text: str) -> str:
|
||||||
|
"""Remove commas that appear before a closing brace/bracket."""
|
||||||
|
import re as _re
|
||||||
|
text = _re.sub(r",\s*}", "}", text)
|
||||||
|
text = _re.sub(r",\s*]", "]", text)
|
||||||
|
return text
|
||||||
@@ -271,12 +271,16 @@ class LoraService(BaseModelService):
|
|||||||
return letters
|
return letters
|
||||||
|
|
||||||
async def get_lora_trigger_words(self, lora_name: str) -> List[str]:
|
async def get_lora_trigger_words(self, lora_name: str) -> List[str]:
|
||||||
"""Get trigger words for a specific LoRA file"""
|
"""Get trigger words for a specific LoRA file.
|
||||||
|
|
||||||
|
Supports both simple names and full-path syntax.
|
||||||
|
"""
|
||||||
cache = await self.scanner.get_cached_data()
|
cache = await self.scanner.get_cached_data()
|
||||||
|
|
||||||
for lora in cache.raw_data:
|
for lora in cache.raw_data:
|
||||||
if lora["file_name"] == lora_name:
|
file_name = lora.get("file_name", "")
|
||||||
civitai_data = lora.get("civitai", {})
|
if file_name == lora_name or lora_name.endswith("/" + file_name) or lora_name.endswith("\\" + file_name):
|
||||||
|
civitai_data = lora.get("civitai") or {}
|
||||||
return civitai_data.get("trainedWords", [])
|
return civitai_data.get("trainedWords", [])
|
||||||
|
|
||||||
return []
|
return []
|
||||||
|
|||||||
@@ -15,6 +15,17 @@ from .service_registry import ServiceRegistry
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_PROVIDER_DISPLAY_NAMES = {
|
||||||
|
"civitai_api": "CivitAI",
|
||||||
|
"civarchive_api": "CivArchive",
|
||||||
|
"sqlite": "Archive DB",
|
||||||
|
}
|
||||||
|
|
||||||
|
_PRESET_PROVIDER_ORDERS = {
|
||||||
|
"civitai_archive_sqlite": ["civitai_api", "civarchive_api", "sqlite"],
|
||||||
|
"civitai_sqlite_archive": ["civitai_api", "sqlite", "civarchive_api"],
|
||||||
|
}
|
||||||
|
|
||||||
async def initialize_metadata_providers():
|
async def initialize_metadata_providers():
|
||||||
"""Initialize and configure all metadata providers based on settings"""
|
"""Initialize and configure all metadata providers based on settings"""
|
||||||
provider_manager = await ModelMetadataProviderManager.get_instance()
|
provider_manager = await ModelMetadataProviderManager.get_instance()
|
||||||
@@ -26,6 +37,8 @@ async def initialize_metadata_providers():
|
|||||||
# Get settings
|
# Get settings
|
||||||
settings_manager = get_settings_manager()
|
settings_manager = get_settings_manager()
|
||||||
enable_archive_db = settings_manager.get('enable_metadata_archive_db', False)
|
enable_archive_db = settings_manager.get('enable_metadata_archive_db', False)
|
||||||
|
enable_civarchive_api = settings_manager.get('enable_civarchive_api', True)
|
||||||
|
provider_order = settings_manager.get('metadata_provider_order', 'civitai_archive_sqlite')
|
||||||
|
|
||||||
providers = []
|
providers = []
|
||||||
|
|
||||||
@@ -59,27 +72,48 @@ async def initialize_metadata_providers():
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Failed to initialize Civitai API metadata provider: {e}")
|
logger.error(f"Failed to initialize Civitai API metadata provider: {e}")
|
||||||
|
|
||||||
# Register CivArchive provider, and all add to fallback providers
|
# Register CivArchive provider when enabled. Civitai API is always
|
||||||
try:
|
# preferred (better metadata); CivArchive mainly recovers metadata for
|
||||||
civarchive_client = await ServiceRegistry.get_civarchive_client()
|
# models deleted from Civitai, so it can be turned off to avoid its long
|
||||||
civarchive_provider = CivArchiveModelMetadataProvider(civarchive_client)
|
# rate-limit windows entirely.
|
||||||
provider_manager.register_provider('civarchive_api', civarchive_provider)
|
if enable_civarchive_api:
|
||||||
providers.append(('civarchive_api', civarchive_provider))
|
try:
|
||||||
logger.debug("CivArchive metadata provider registered (also included in fallback)")
|
civarchive_client = await ServiceRegistry.get_civarchive_client()
|
||||||
except Exception as e:
|
civarchive_provider = CivArchiveModelMetadataProvider(civarchive_client)
|
||||||
logger.error(f"Failed to initialize CivArchive metadata provider: {e}")
|
provider_manager.register_provider('civarchive_api', civarchive_provider)
|
||||||
|
providers.append(('civarchive_api', civarchive_provider))
|
||||||
|
logger.debug("CivArchive metadata provider registered (also included in fallback)")
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Failed to initialize CivArchive metadata provider: {e}")
|
||||||
|
else:
|
||||||
|
logger.debug("CivArchive metadata provider disabled by setting 'enable_civarchive_api'")
|
||||||
|
|
||||||
|
# Preset fallback orderings (see module-level _PRESET_PROVIDER_ORDERS).
|
||||||
|
# civitai_api is always first (better metadata); the remaining providers
|
||||||
|
# are arranged by the configured preset. Providers that are not
|
||||||
|
# registered (disabled/unavailable) are simply skipped, so each preset
|
||||||
|
# degrades gracefully.
|
||||||
|
desired_order = _PRESET_PROVIDER_ORDERS.get(
|
||||||
|
provider_order, _PRESET_PROVIDER_ORDERS["civitai_archive_sqlite"]
|
||||||
|
)
|
||||||
|
|
||||||
# Set up fallback provider based on available providers
|
# Set up fallback provider based on available providers
|
||||||
if len(providers) > 1:
|
if len(providers) > 1:
|
||||||
# Always use Civitai API (it has better metadata), then CivArchive API, then Archive DB
|
|
||||||
ordered_providers: list[tuple[str, ModelMetadataProvider]] = []
|
ordered_providers: list[tuple[str, ModelMetadataProvider]] = []
|
||||||
ordered_providers.extend([p for p in providers if p[0] == 'civitai_api'])
|
for name in desired_order:
|
||||||
ordered_providers.extend([p for p in providers if p[0] == 'civarchive_api'])
|
ordered_providers.extend([p for p in providers if p[0] == name])
|
||||||
ordered_providers.extend([p for p in providers if p[0] == 'sqlite'])
|
# Include any provider not covered by the preset (defensive) at the end
|
||||||
|
for p in providers:
|
||||||
|
if p not in ordered_providers:
|
||||||
|
ordered_providers.append(p)
|
||||||
|
|
||||||
if ordered_providers:
|
if ordered_providers:
|
||||||
fallback_provider = FallbackMetadataProvider(ordered_providers)
|
fallback_provider = FallbackMetadataProvider(ordered_providers)
|
||||||
provider_manager.register_provider('fallback', fallback_provider, is_default=True)
|
provider_manager.register_provider('fallback', fallback_provider, is_default=True)
|
||||||
|
logger.debug(
|
||||||
|
"Metadata fallback provider order: %s",
|
||||||
|
", ".join(name for name, _ in ordered_providers),
|
||||||
|
)
|
||||||
elif len(providers) == 1:
|
elif len(providers) == 1:
|
||||||
# Only one provider available, set it as default
|
# Only one provider available, set it as default
|
||||||
provider_name, provider = providers[0]
|
provider_name, provider = providers[0]
|
||||||
@@ -96,11 +130,30 @@ async def update_metadata_providers():
|
|||||||
# Get current settings
|
# Get current settings
|
||||||
settings_manager = get_settings_manager()
|
settings_manager = get_settings_manager()
|
||||||
enable_archive_db = settings_manager.get('enable_metadata_archive_db', False)
|
enable_archive_db = settings_manager.get('enable_metadata_archive_db', False)
|
||||||
|
enable_civarchive_api = settings_manager.get('enable_civarchive_api', True)
|
||||||
|
provider_order = settings_manager.get('metadata_provider_order', 'civitai_archive_sqlite')
|
||||||
|
|
||||||
# Reinitialize all providers with new settings
|
# Reinitialize all providers with new settings
|
||||||
provider_manager = await initialize_metadata_providers()
|
provider_manager = await initialize_metadata_providers()
|
||||||
|
|
||||||
logger.info(f"Updated metadata providers, archive_db enabled: {enable_archive_db}")
|
# Build effective provider chain for logging (use actually-registered
|
||||||
|
# providers, not just settings, so a failed init is reflected correctly)
|
||||||
|
registered = set(provider_manager.providers.keys())
|
||||||
|
desired = _PRESET_PROVIDER_ORDERS.get(
|
||||||
|
provider_order, _PRESET_PROVIDER_ORDERS["civitai_archive_sqlite"]
|
||||||
|
)
|
||||||
|
chain = " → ".join(
|
||||||
|
_PROVIDER_DISPLAY_NAMES[p]
|
||||||
|
for p in desired
|
||||||
|
if p in registered and p in _PROVIDER_DISPLAY_NAMES
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Updated metadata providers: archive_db=%s, civarchive_api=%s, chain=%s",
|
||||||
|
enable_archive_db,
|
||||||
|
enable_civarchive_api,
|
||||||
|
chain,
|
||||||
|
)
|
||||||
return provider_manager
|
return provider_manager
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Failed to update metadata providers: {e}")
|
logger.error(f"Failed to update metadata providers: {e}")
|
||||||
|
|||||||
@@ -209,7 +209,21 @@ class MetadataSyncService:
|
|||||||
error_msg = "CivitAI model is deleted and no archive provider is available"
|
error_msg = "CivitAI model is deleted and no archive provider is available"
|
||||||
return False, error_msg
|
return False, error_msg
|
||||||
else:
|
else:
|
||||||
provider_attempts.append((None, await self._get_default_provider()))
|
is_hf_source = bool(model_data.get("hf_url"))
|
||||||
|
if is_hf_source:
|
||||||
|
# HF-sourced model: only check CivitAI API directly.
|
||||||
|
# CivArchive is almost guaranteed to have no record, and
|
||||||
|
# hitting it wastes rate-limit budget.
|
||||||
|
# Use a distinct provider name ("civitai_api" not None) so
|
||||||
|
# downstream code does NOT interpret a "Model not found"
|
||||||
|
# response as civitai_api_not_found — which would mark the
|
||||||
|
# model civitai_deleted=True when it was never on CivitAI.
|
||||||
|
try:
|
||||||
|
provider_attempts.append(("civitai_api", await self._get_provider("civitai_api")))
|
||||||
|
except Exception as exc: # pragma: no cover - provider resolution fault
|
||||||
|
logger.debug("Unable to resolve civitai_api provider: %s", exc)
|
||||||
|
if not provider_attempts:
|
||||||
|
provider_attempts.append((None, await self._get_default_provider()))
|
||||||
|
|
||||||
civitai_metadata: Optional[Dict[str, Any]] = None
|
civitai_metadata: Optional[Dict[str, Any]] = None
|
||||||
metadata_provider: Optional[MetadataProviderProtocol] = None
|
metadata_provider: Optional[MetadataProviderProtocol] = None
|
||||||
|
|||||||
@@ -338,3 +338,24 @@ class ModelCache:
|
|||||||
return False # Model not found
|
return False # Model not found
|
||||||
|
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
async def clear_preview_by_path(self, preview_file_path: str) -> int:
|
||||||
|
"""Clear ``preview_url`` for every cached entry referencing a file path.
|
||||||
|
|
||||||
|
When a preview file has been deleted from disk, this removes its
|
||||||
|
reference from all matching cache entries so the next list-API
|
||||||
|
response returns an empty ``preview_url`` instead of a stale URL
|
||||||
|
that produces 404s.
|
||||||
|
|
||||||
|
Returns the number of entries that were updated.
|
||||||
|
"""
|
||||||
|
normalized = preview_file_path.replace("\\", "/")
|
||||||
|
cleared = 0
|
||||||
|
async with self._lock:
|
||||||
|
for item in self.raw_data:
|
||||||
|
cached_url = item.get("preview_url", "")
|
||||||
|
if cached_url.replace("\\", "/") == normalized:
|
||||||
|
item["preview_url"] = ""
|
||||||
|
item["preview_nsfw_level"] = 0
|
||||||
|
cleared += 1
|
||||||
|
return cleared
|
||||||
@@ -8,6 +8,7 @@ from abc import ABC, abstractmethod
|
|||||||
from ..utils.utils import calculate_relative_path_for_model, remove_empty_dirs
|
from ..utils.utils import calculate_relative_path_for_model, remove_empty_dirs
|
||||||
from ..utils.constants import AUTO_ORGANIZE_BATCH_SIZE
|
from ..utils.constants import AUTO_ORGANIZE_BATCH_SIZE
|
||||||
from ..services.settings_manager import get_settings_manager
|
from ..services.settings_manager import get_settings_manager
|
||||||
|
from ..services.model_lifecycle_service import _require_path_in_library_roots
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -493,6 +494,9 @@ class ModelMoveService:
|
|||||||
Dictionary with move result
|
Dictionary with move result
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
|
_require_path_in_library_roots(file_path, self.scanner, label="Source path")
|
||||||
|
_require_path_in_library_roots(target_path, self.scanner, label="Target path")
|
||||||
|
|
||||||
if use_default_paths:
|
if use_default_paths:
|
||||||
# Find the model in cache to get metadata
|
# Find the model in cache to get metadata
|
||||||
cache = await self.scanner.get_cached_data()
|
cache = await self.scanner.get_cached_data()
|
||||||
|
|||||||
@@ -48,6 +48,36 @@ async def delete_model_artifacts(
|
|||||||
return deleted
|
return deleted
|
||||||
|
|
||||||
|
|
||||||
|
def _require_path_in_library_roots(file_path: str, scanner, *, label: str = "path") -> None:
|
||||||
|
"""Raise ``ValueError`` if *file_path* is not inside a configured model root.
|
||||||
|
|
||||||
|
Uses ``os.path.abspath()`` (NOT ``realpath``) to resolve ``..`` and ``.``
|
||||||
|
while preserving symlinks — this keeps the check in business-path space.
|
||||||
|
Skips when the scanner does not expose ``get_model_roots`` or the list
|
||||||
|
is empty.
|
||||||
|
"""
|
||||||
|
|
||||||
|
roots = None
|
||||||
|
if hasattr(scanner, "get_model_roots"):
|
||||||
|
try:
|
||||||
|
roots = scanner.get_model_roots()
|
||||||
|
except NotImplementedError:
|
||||||
|
roots = None
|
||||||
|
if not roots:
|
||||||
|
return
|
||||||
|
|
||||||
|
resolved = os.path.abspath(os.path.normpath(file_path))
|
||||||
|
|
||||||
|
for root in roots:
|
||||||
|
root_resolved = os.path.abspath(os.path.normpath(root))
|
||||||
|
if resolved == root_resolved or resolved.startswith(root_resolved + os.sep):
|
||||||
|
return
|
||||||
|
|
||||||
|
raise ValueError(
|
||||||
|
f"{label} '{file_path}' is outside configured library directories"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class ModelLifecycleService:
|
class ModelLifecycleService:
|
||||||
"""Co-ordinate destructive and mutating model operations."""
|
"""Co-ordinate destructive and mutating model operations."""
|
||||||
|
|
||||||
@@ -74,6 +104,8 @@ class ModelLifecycleService:
|
|||||||
if not file_path:
|
if not file_path:
|
||||||
raise ValueError("Model path is required")
|
raise ValueError("Model path is required")
|
||||||
|
|
||||||
|
_require_path_in_library_roots(file_path, self._scanner, label="File path")
|
||||||
|
|
||||||
cache = await self._scanner.get_cached_data()
|
cache = await self._scanner.get_cached_data()
|
||||||
|
|
||||||
cached_entry = None
|
cached_entry = None
|
||||||
@@ -182,6 +214,8 @@ class ModelLifecycleService:
|
|||||||
if not file_path:
|
if not file_path:
|
||||||
raise ValueError("Model path is required")
|
raise ValueError("Model path is required")
|
||||||
|
|
||||||
|
_require_path_in_library_roots(file_path, self._scanner, label="File path")
|
||||||
|
|
||||||
metadata_path = os.path.splitext(file_path)[0] + ".metadata.json"
|
metadata_path = os.path.splitext(file_path)[0] + ".metadata.json"
|
||||||
metadata = await self._metadata_loader(metadata_path)
|
metadata = await self._metadata_loader(metadata_path)
|
||||||
metadata["exclude"] = True
|
metadata["exclude"] = True
|
||||||
@@ -229,6 +263,8 @@ class ModelLifecycleService:
|
|||||||
if not file_path:
|
if not file_path:
|
||||||
raise ValueError("Model path is required")
|
raise ValueError("Model path is required")
|
||||||
|
|
||||||
|
_require_path_in_library_roots(file_path, self._scanner, label="File path")
|
||||||
|
|
||||||
if not os.path.exists(file_path):
|
if not os.path.exists(file_path):
|
||||||
raise ValueError("Model file does not exist")
|
raise ValueError("Model file does not exist")
|
||||||
|
|
||||||
@@ -270,6 +306,9 @@ class ModelLifecycleService:
|
|||||||
if not file_paths:
|
if not file_paths:
|
||||||
raise ValueError("No file paths provided for deletion")
|
raise ValueError("No file paths provided for deletion")
|
||||||
|
|
||||||
|
for path in file_paths:
|
||||||
|
_require_path_in_library_roots(path, self._scanner, label="File path")
|
||||||
|
|
||||||
return await self._scanner.bulk_delete_models(file_paths)
|
return await self._scanner.bulk_delete_models(file_paths)
|
||||||
|
|
||||||
async def rename_model(
|
async def rename_model(
|
||||||
@@ -280,6 +319,8 @@ class ModelLifecycleService:
|
|||||||
if not file_path or not new_file_name:
|
if not file_path or not new_file_name:
|
||||||
raise ValueError("File path and new file name are required")
|
raise ValueError("File path and new file name are required")
|
||||||
|
|
||||||
|
_require_path_in_library_roots(file_path, self._scanner, label="File path")
|
||||||
|
|
||||||
invalid_chars = {"/", "\\", ":", "*", "?", '"', "<", ">", "|"}
|
invalid_chars = {"/", "\\", ":", "*", "?", '"', "<", ">", "|"}
|
||||||
if any(char in new_file_name for char in invalid_chars):
|
if any(char in new_file_name for char in invalid_chars):
|
||||||
raise ValueError("Invalid characters in file name")
|
raise ValueError("Invalid characters in file name")
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ from ..utils.metadata_manager import MetadataManager
|
|||||||
from ..utils.civitai_utils import resolve_license_info
|
from ..utils.civitai_utils import resolve_license_info
|
||||||
from .model_cache import ModelCache
|
from .model_cache import ModelCache
|
||||||
from .model_hash_index import ModelHashIndex
|
from .model_hash_index import ModelHashIndex
|
||||||
from .model_lifecycle_service import delete_model_artifacts
|
from .model_lifecycle_service import delete_model_artifacts, _require_path_in_library_roots
|
||||||
from .service_registry import ServiceRegistry
|
from .service_registry import ServiceRegistry
|
||||||
from .websocket_manager import ws_manager
|
from .websocket_manager import ws_manager
|
||||||
from .persistent_model_cache import get_persistent_cache
|
from .persistent_model_cache import get_persistent_cache
|
||||||
@@ -227,6 +227,11 @@ class ModelScanner:
|
|||||||
|
|
||||||
entry: Dict[str, Any] = {
|
entry: Dict[str, Any] = {
|
||||||
'file_path': normalized_path,
|
'file_path': normalized_path,
|
||||||
|
# file_name is always stored WITHOUT extension (e.g. "OWSMianne_ANIMA_V1",
|
||||||
|
# not "OWSMianne_ANIMA_V1.safetensors"). All upstream population points
|
||||||
|
# (MetadataManager, from_civitai_info, download manager, etc.) strip the
|
||||||
|
# extension via os.path.splitext before writing. Code consuming this field
|
||||||
|
# should match against names that are likewise extension-free.
|
||||||
'file_name': get_value('file_name', '') or '',
|
'file_name': get_value('file_name', '') or '',
|
||||||
'model_name': get_value('model_name', '') or '',
|
'model_name': get_value('model_name', '') or '',
|
||||||
'folder': normalized_folder,
|
'folder': normalized_folder,
|
||||||
@@ -922,6 +927,25 @@ class ModelScanner:
|
|||||||
# Update cache data
|
# Update cache data
|
||||||
self._cache.raw_data = [item for item in self._cache.raw_data if item['file_path'] not in missing_files]
|
self._cache.raw_data = [item for item in self._cache.raw_data if item['file_path'] not in missing_files]
|
||||||
|
|
||||||
|
dedup_removed = 0
|
||||||
|
seen_paths: set = set()
|
||||||
|
deduped: list = []
|
||||||
|
for item in reversed(self._cache.raw_data):
|
||||||
|
path = item.get('file_path', '')
|
||||||
|
if path not in seen_paths:
|
||||||
|
seen_paths.add(path)
|
||||||
|
deduped.append(item)
|
||||||
|
else:
|
||||||
|
for tag in item.get('tags', []):
|
||||||
|
if tag in self._tags_count:
|
||||||
|
self._tags_count[tag] = max(0, self._tags_count[tag] - 1)
|
||||||
|
if self._tags_count[tag] == 0:
|
||||||
|
del self._tags_count[tag]
|
||||||
|
dedup_removed += 1
|
||||||
|
if dedup_removed > 0:
|
||||||
|
self._cache.raw_data = list(reversed(deduped))
|
||||||
|
total_removed += dedup_removed
|
||||||
|
|
||||||
# Resort cache if changes were made
|
# Resort cache if changes were made
|
||||||
if total_added > 0 or total_removed > 0:
|
if total_added > 0 or total_removed > 0:
|
||||||
# Update folders list
|
# Update folders list
|
||||||
@@ -1347,18 +1371,25 @@ class ModelScanner:
|
|||||||
# Update folder in metadata
|
# Update folder in metadata
|
||||||
metadata_dict['folder'] = folder
|
metadata_dict['folder'] = folder
|
||||||
|
|
||||||
# Add to cache
|
file_path = metadata_dict.get('file_path', '')
|
||||||
|
if file_path:
|
||||||
|
old_entries = [item for item in self._cache.raw_data if item.get('file_path') == file_path]
|
||||||
|
for old_entry in old_entries:
|
||||||
|
for tag in old_entry.get('tags', []):
|
||||||
|
if tag in self._tags_count:
|
||||||
|
self._tags_count[tag] = max(0, self._tags_count[tag] - 1)
|
||||||
|
if self._tags_count[tag] == 0:
|
||||||
|
del self._tags_count[tag]
|
||||||
|
self._hash_index.remove_by_path(file_path)
|
||||||
|
self._cache.raw_data = [item for item in self._cache.raw_data if item.get('file_path') != file_path]
|
||||||
|
|
||||||
|
for tag in metadata_dict.get('tags', []):
|
||||||
|
self._tags_count[tag] = self._tags_count.get(tag, 0) + 1
|
||||||
|
|
||||||
self._cache.raw_data.append(metadata_dict)
|
self._cache.raw_data.append(metadata_dict)
|
||||||
self._cache.add_to_version_index(metadata_dict)
|
|
||||||
|
|
||||||
# Resort cache data
|
|
||||||
await self._cache.resort()
|
await self._cache.resort()
|
||||||
|
|
||||||
# Update folders list
|
|
||||||
all_folders = set(self._cache.folders)
|
|
||||||
all_folders.add(folder)
|
|
||||||
self._cache.folders = sorted(list(all_folders), key=lambda x: x.lower())
|
|
||||||
|
|
||||||
# Update the hash index
|
# Update the hash index
|
||||||
self._hash_index.add_entry(metadata_dict['sha256'], metadata_dict['file_path'])
|
self._hash_index.add_entry(metadata_dict['sha256'], metadata_dict['file_path'])
|
||||||
await self._persist_current_cache()
|
await self._persist_current_cache()
|
||||||
@@ -1390,6 +1421,9 @@ class ModelScanner:
|
|||||||
base_name = os.path.splitext(os.path.basename(source_path))[0]
|
base_name = os.path.splitext(os.path.basename(source_path))[0]
|
||||||
source_dir = os.path.dirname(source_path)
|
source_dir = os.path.dirname(source_path)
|
||||||
|
|
||||||
|
_require_path_in_library_roots(source_path, self, label="Source path")
|
||||||
|
_require_path_in_library_roots(target_path, self, label="Target path")
|
||||||
|
|
||||||
os.makedirs(target_path, exist_ok=True)
|
os.makedirs(target_path, exist_ok=True)
|
||||||
|
|
||||||
def get_source_hash():
|
def get_source_hash():
|
||||||
@@ -1561,6 +1595,218 @@ class ModelScanner:
|
|||||||
|
|
||||||
return cache_entry if metadata else True
|
return cache_entry if metadata else True
|
||||||
|
|
||||||
|
async def sync_cache_from_metadata(
|
||||||
|
self, file_path: str, metadata_dict: Dict[str, Any]
|
||||||
|
) -> bool:
|
||||||
|
"""Opportunistically sync in-memory and persistent caches from metadata.
|
||||||
|
|
||||||
|
Builds a prospective cache entry from *metadata_dict* (deserialized
|
||||||
|
``.metadata.json`` content) and compares it against the current cache
|
||||||
|
entry. When the two are already identical this method returns
|
||||||
|
``False`` without touching anything — avoiding the overhead of
|
||||||
|
``update_single_model_cache``, which always removes and re-inserts
|
||||||
|
the entry, triggers a full resort, and persists via the heavyweight
|
||||||
|
``save_cache()``.
|
||||||
|
|
||||||
|
When differences are detected the update is applied **in-place** with
|
||||||
|
targeted operations:
|
||||||
|
|
||||||
|
* The existing ``raw_data`` entry is modified rather than removed and
|
||||||
|
re-appended (O(1) instead of O(n)).
|
||||||
|
* Tag counts and the hash index are updated incrementally.
|
||||||
|
* The version index is rebuilt only for the affected entry.
|
||||||
|
* ``resort()`` is called **only** when a sort-relevant field changed
|
||||||
|
(``model_name`` / ``file_name`` for name-sort, ``modified`` for
|
||||||
|
date-sort, ``size`` for size-sort).
|
||||||
|
* The persistent (SQLite) cache receives a targeted single-row update
|
||||||
|
via :meth:`PersistentModelCache.update_single_model` rather than a
|
||||||
|
full-table ``save_cache()``.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``True`` if any cache update was performed, ``False`` if the
|
||||||
|
caches were already in sync.
|
||||||
|
|
||||||
|
.. note::
|
||||||
|
|
||||||
|
This is a **best-effort** operation. Failures are logged but
|
||||||
|
never propagated — callers should fire-and-forget via
|
||||||
|
:func:`asyncio.create_task`.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
return await self._sync_cache_from_metadata_impl(
|
||||||
|
file_path, metadata_dict
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.warning(
|
||||||
|
"sync_cache_from_metadata failed for %s",
|
||||||
|
file_path,
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def _sync_cache_from_metadata_impl(
|
||||||
|
self, file_path: str, metadata_dict: Dict[str, Any]
|
||||||
|
) -> bool:
|
||||||
|
cache = await self.get_cached_data()
|
||||||
|
|
||||||
|
# Locate the existing cache entry -----------------------------------
|
||||||
|
existing_idx: Optional[int] = None
|
||||||
|
existing_entry: Optional[Dict[str, Any]] = None
|
||||||
|
for i, item in enumerate(cache.raw_data):
|
||||||
|
if item.get("file_path") == file_path:
|
||||||
|
existing_entry = item
|
||||||
|
existing_idx = i
|
||||||
|
break
|
||||||
|
|
||||||
|
# Build the desired entry from metadata ------------------------------
|
||||||
|
folder_value = (
|
||||||
|
existing_entry.get("folder", "")
|
||||||
|
if existing_entry
|
||||||
|
else self._calculate_folder(file_path)
|
||||||
|
)
|
||||||
|
desired_entry = self._build_cache_entry(
|
||||||
|
metadata_dict,
|
||||||
|
folder=folder_value,
|
||||||
|
file_path_override=file_path,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Ensure sha256 is populated (defensive — metadata should have it)
|
||||||
|
if (
|
||||||
|
not desired_entry.get("sha256")
|
||||||
|
and file_path
|
||||||
|
and os.path.exists(file_path)
|
||||||
|
):
|
||||||
|
try:
|
||||||
|
sha256 = await calculate_sha256(file_path)
|
||||||
|
if sha256:
|
||||||
|
desired_entry["sha256"] = sha256.lower()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Not in cache at all — delegate to the full update path ------------
|
||||||
|
if existing_entry is None:
|
||||||
|
result = await self.update_single_model_cache(
|
||||||
|
file_path, file_path, metadata_dict
|
||||||
|
)
|
||||||
|
return bool(result)
|
||||||
|
|
||||||
|
# Compare — skip everything if already in sync -----------------------
|
||||||
|
if not self._cache_entries_differ(existing_entry, desired_entry):
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Re-validate: the cache may have been replaced concurrently
|
||||||
|
# (e.g. by _apply_scan_result). Use identity check, not equality,
|
||||||
|
# so we detect when the raw_data list was swapped out from under us.
|
||||||
|
if self._cache is None or not any(
|
||||||
|
item is existing_entry for item in self._cache.raw_data
|
||||||
|
):
|
||||||
|
return False
|
||||||
|
|
||||||
|
# ---- Differences detected: apply targeted, in-place updates --------
|
||||||
|
|
||||||
|
# Snapshot old values for delta computations
|
||||||
|
old_tags = list(existing_entry.get("tags") or [])
|
||||||
|
old_sha256: str = existing_entry.get("sha256", "") or ""
|
||||||
|
old_model_name: str = existing_entry.get("model_name", "") or ""
|
||||||
|
old_file_name: str = existing_entry.get("file_name", "") or ""
|
||||||
|
old_modified: float = float(existing_entry.get("modified", 0.0) or 0.0)
|
||||||
|
old_size: int = int(existing_entry.get("size", 0) or 0)
|
||||||
|
old_civitai = existing_entry.get("civitai")
|
||||||
|
|
||||||
|
# ---- In-place update of the cache entry ----
|
||||||
|
existing_entry.clear()
|
||||||
|
existing_entry.update(desired_entry)
|
||||||
|
|
||||||
|
# ---- Incremental tag count update ----
|
||||||
|
new_tags: set = set(desired_entry.get("tags") or [])
|
||||||
|
old_tag_set: set = set(old_tags)
|
||||||
|
for tag in old_tag_set - new_tags:
|
||||||
|
current = self._tags_count.get(tag, 0)
|
||||||
|
if current <= 1:
|
||||||
|
self._tags_count.pop(tag, None)
|
||||||
|
else:
|
||||||
|
self._tags_count[tag] = current - 1
|
||||||
|
for tag in new_tags - old_tag_set:
|
||||||
|
self._tags_count[tag] = self._tags_count.get(tag, 0) + 1
|
||||||
|
|
||||||
|
# ---- Incremental hash index update ----
|
||||||
|
new_sha = (desired_entry.get("sha256", "") or "").lower()
|
||||||
|
old_sha = (old_sha256 or "").lower()
|
||||||
|
if new_sha != old_sha:
|
||||||
|
if old_sha:
|
||||||
|
self._hash_index.remove_by_path(file_path)
|
||||||
|
if new_sha:
|
||||||
|
self._hash_index.add_entry(new_sha, file_path)
|
||||||
|
|
||||||
|
# ---- Incremental version index update ----
|
||||||
|
new_civitai = desired_entry.get("civitai")
|
||||||
|
if old_civitai != new_civitai:
|
||||||
|
temp_old = {
|
||||||
|
"file_path": file_path,
|
||||||
|
"file_name": old_file_name,
|
||||||
|
"civitai": old_civitai,
|
||||||
|
}
|
||||||
|
cache.remove_from_version_index(temp_old)
|
||||||
|
cache.add_to_version_index(existing_entry)
|
||||||
|
|
||||||
|
# ---- Conditional resort (only when sort-key fields changed) ----
|
||||||
|
need_resort = False
|
||||||
|
_last = cache._last_sort
|
||||||
|
sort_key: Optional[str] = _last[0] if _last != (None, None) else None
|
||||||
|
if sort_key == "name":
|
||||||
|
if (
|
||||||
|
old_model_name != desired_entry.get("model_name", "")
|
||||||
|
or old_file_name != desired_entry.get("file_name", "")
|
||||||
|
):
|
||||||
|
need_resort = True
|
||||||
|
elif sort_key == "date":
|
||||||
|
if old_modified != float(desired_entry.get("modified", 0.0) or 0.0):
|
||||||
|
need_resort = True
|
||||||
|
elif sort_key == "size":
|
||||||
|
if old_size != int(desired_entry.get("size", 0) or 0):
|
||||||
|
need_resort = True
|
||||||
|
|
||||||
|
if need_resort:
|
||||||
|
await cache.resort()
|
||||||
|
|
||||||
|
# ---- Targeted SQL update (single row, not full save_cache) ----
|
||||||
|
persistent = getattr(self, "_persistent_cache", None)
|
||||||
|
if persistent is not None:
|
||||||
|
old_item_for_sql: Dict[str, Any] = {
|
||||||
|
"file_path": file_path,
|
||||||
|
"tags": old_tags,
|
||||||
|
"sha256": old_sha256,
|
||||||
|
}
|
||||||
|
await asyncio.get_event_loop().run_in_executor(
|
||||||
|
None,
|
||||||
|
persistent.update_single_model,
|
||||||
|
self.model_type,
|
||||||
|
desired_entry,
|
||||||
|
old_item_for_sql,
|
||||||
|
)
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _cache_entries_differ(a: Dict[str, Any], b: Dict[str, Any]) -> bool:
|
||||||
|
"""Return ``True`` when two cache-entry dicts differ in any field.
|
||||||
|
|
||||||
|
Tag lists are compared order-insensitively; all other keys use
|
||||||
|
standard equality.
|
||||||
|
"""
|
||||||
|
a_tags = sorted(a.get("tags") or [])
|
||||||
|
b_tags = sorted(b.get("tags") or [])
|
||||||
|
if a_tags != b_tags:
|
||||||
|
return True
|
||||||
|
|
||||||
|
all_keys = set(a.keys()) | set(b.keys())
|
||||||
|
for key in all_keys:
|
||||||
|
if key == "tags":
|
||||||
|
continue
|
||||||
|
if a.get(key) != b.get(key):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
def has_hash(self, sha256: str) -> bool:
|
def has_hash(self, sha256: str) -> bool:
|
||||||
"""Check if a model with given hash exists"""
|
"""Check if a model with given hash exists"""
|
||||||
return self._hash_index.has_hash(sha256.lower())
|
return self._hash_index.has_hash(sha256.lower())
|
||||||
@@ -1614,6 +1860,31 @@ class ModelScanner:
|
|||||||
return sorted_tags
|
return sorted_tags
|
||||||
return sorted_tags[:limit]
|
return sorted_tags[:limit]
|
||||||
|
|
||||||
|
async def search_tags(
|
||||||
|
self, query: str, limit: int = 50
|
||||||
|
) -> List[Dict[str, any]]:
|
||||||
|
"""Search tags by case-insensitive substring match, sorted by count.
|
||||||
|
|
||||||
|
If query is empty, behaves like get_top_tags (returns top ``limit``
|
||||||
|
tags). If limit is 0, all matching tags are returned.
|
||||||
|
"""
|
||||||
|
await self.get_cached_data()
|
||||||
|
|
||||||
|
normalized_query = (query or "").strip().lower()
|
||||||
|
if not normalized_query:
|
||||||
|
return await self.get_top_tags(limit if limit > 0 else 20)
|
||||||
|
|
||||||
|
matched = [
|
||||||
|
{"tag": tag, "count": count}
|
||||||
|
for tag, count in self._tags_count.items()
|
||||||
|
if normalized_query in tag.lower()
|
||||||
|
]
|
||||||
|
matched.sort(key=lambda x: x["count"], reverse=True)
|
||||||
|
|
||||||
|
if limit == 0:
|
||||||
|
return matched
|
||||||
|
return matched[:limit]
|
||||||
|
|
||||||
async def get_base_models(self, limit: int = 20) -> List[Dict[str, any]]:
|
async def get_base_models(self, limit: int = 20) -> List[Dict[str, any]]:
|
||||||
"""Get base models sorted by count. If limit is 0, return all."""
|
"""Get base models sorted by count. If limit is 0, return all."""
|
||||||
cache = await self.get_cached_data()
|
cache = await self.get_cached_data()
|
||||||
@@ -1729,6 +2000,8 @@ class ModelScanner:
|
|||||||
break
|
break
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
_require_path_in_library_roots(file_path, self, label="File path")
|
||||||
|
|
||||||
target_dir = os.path.dirname(file_path)
|
target_dir = os.path.dirname(file_path)
|
||||||
base_name = os.path.basename(file_path)
|
base_name = os.path.basename(file_path)
|
||||||
file_name, main_extension = os.path.splitext(base_name)
|
file_name, main_extension = os.path.splitext(base_name)
|
||||||
|
|||||||
@@ -587,6 +587,95 @@ class PersistentModelCache:
|
|||||||
placeholders = ", ".join(["?"] * len(self._MODEL_COLUMNS))
|
placeholders = ", ".join(["?"] * len(self._MODEL_COLUMNS))
|
||||||
return f"INSERT INTO models ({columns}) VALUES ({placeholders})"
|
return f"INSERT INTO models ({columns}) VALUES ({placeholders})"
|
||||||
|
|
||||||
|
def update_single_model(
|
||||||
|
self,
|
||||||
|
model_type: str,
|
||||||
|
new_item: Dict,
|
||||||
|
old_item: Optional[Dict] = None,
|
||||||
|
) -> None:
|
||||||
|
"""Update a single model row in the persistent cache.
|
||||||
|
|
||||||
|
A lightweight alternative to :meth:`save_cache` that performs a targeted
|
||||||
|
DELETE + INSERT for the model row and computes incremental tag / hash-index
|
||||||
|
deltas from *old_item*. When *old_item* is omitted the previous tags and
|
||||||
|
hash are not cleaned up (callers should only omit it for brand-new entries).
|
||||||
|
|
||||||
|
All operations run inside a single transaction so readers see a consistent
|
||||||
|
view.
|
||||||
|
"""
|
||||||
|
if not self.is_enabled():
|
||||||
|
return
|
||||||
|
if not self._schema_initialized:
|
||||||
|
self._initialize_schema()
|
||||||
|
if not self._schema_initialized:
|
||||||
|
return
|
||||||
|
|
||||||
|
file_path: Optional[str] = new_item.get("file_path")
|
||||||
|
if not file_path:
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
with self._db_lock:
|
||||||
|
conn = self._connect()
|
||||||
|
try:
|
||||||
|
conn.execute("PRAGMA foreign_keys = ON")
|
||||||
|
conn.execute("BEGIN")
|
||||||
|
|
||||||
|
# --- model row (DELETE + INSERT = upsert) ---
|
||||||
|
conn.execute(
|
||||||
|
"DELETE FROM models WHERE model_type = ? AND file_path = ?",
|
||||||
|
(model_type, file_path),
|
||||||
|
)
|
||||||
|
row = self._prepare_model_row(model_type, new_item)
|
||||||
|
conn.execute(self._insert_model_sql(), row)
|
||||||
|
|
||||||
|
# --- tags ---
|
||||||
|
new_tags: set = set(new_item.get("tags") or [])
|
||||||
|
old_tags: set = set(old_item.get("tags") or []) if old_item else set()
|
||||||
|
tags_to_delete = old_tags - new_tags
|
||||||
|
tags_to_insert = new_tags - old_tags
|
||||||
|
|
||||||
|
if tags_to_delete:
|
||||||
|
conn.executemany(
|
||||||
|
"DELETE FROM model_tags WHERE model_type = ? AND file_path = ? AND tag = ?",
|
||||||
|
[(model_type, file_path, t) for t in tags_to_delete],
|
||||||
|
)
|
||||||
|
if tags_to_insert:
|
||||||
|
conn.executemany(
|
||||||
|
"INSERT INTO model_tags (model_type, file_path, tag) VALUES (?, ?, ?)",
|
||||||
|
[(model_type, file_path, t) for t in tags_to_insert],
|
||||||
|
)
|
||||||
|
|
||||||
|
# --- hash_index ---
|
||||||
|
new_sha: Optional[str] = (new_item.get("sha256") or "").lower() or None
|
||||||
|
old_sha: Optional[str] = (
|
||||||
|
(old_item.get("sha256") or "").lower() or None
|
||||||
|
) if old_item else None
|
||||||
|
if new_sha != old_sha:
|
||||||
|
if old_sha:
|
||||||
|
conn.execute(
|
||||||
|
"DELETE FROM hash_index WHERE model_type = ? AND sha256 = ? AND file_path = ?",
|
||||||
|
(model_type, old_sha, file_path),
|
||||||
|
)
|
||||||
|
if new_sha:
|
||||||
|
conn.execute(
|
||||||
|
"INSERT OR IGNORE INTO hash_index (model_type, sha256, file_path) VALUES (?, ?, ?)",
|
||||||
|
(model_type, new_sha, file_path),
|
||||||
|
)
|
||||||
|
|
||||||
|
conn.execute("COMMIT")
|
||||||
|
except Exception:
|
||||||
|
conn.execute("ROLLBACK")
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
conn.close()
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to update single model in persistent cache (%s): %s",
|
||||||
|
file_path,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
|
||||||
def _load_tags(self, conn: sqlite3.Connection, model_type: str) -> Dict[str, List[str]]:
|
def _load_tags(self, conn: sqlite3.Connection, model_type: str) -> Dict[str, List[str]]:
|
||||||
tag_rows = conn.execute(
|
tag_rows = conn.execute(
|
||||||
"SELECT file_path, tag FROM model_tags WHERE model_type = ?",
|
"SELECT file_path, tag FROM model_tags WHERE model_type = ?",
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
from typing import Iterable, List, Dict, Optional
|
from typing import Iterable, List, Dict, Optional
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from operator import itemgetter
|
|
||||||
from natsort import natsorted
|
from natsort import natsorted
|
||||||
|
|
||||||
|
|
||||||
@@ -149,5 +148,10 @@ class RecipeCache:
|
|||||||
)
|
)
|
||||||
if not name_only:
|
if not name_only:
|
||||||
self.sorted_by_date = sorted(
|
self.sorted_by_date = sorted(
|
||||||
self.raw_data, key=itemgetter("created_date", "file_path"), reverse=True
|
self.raw_data,
|
||||||
|
key=lambda x: (
|
||||||
|
x.get("modified", x.get("created_date", 0)),
|
||||||
|
x.get("file_path", ""),
|
||||||
|
),
|
||||||
|
reverse=True,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ from .checkpoint_scanner import CheckpointScanner
|
|||||||
from .settings_manager import get_settings_manager
|
from .settings_manager import get_settings_manager
|
||||||
from .recipes.errors import RecipeNotFoundError
|
from .recipes.errors import RecipeNotFoundError
|
||||||
from ..utils.civitai_utils import extract_civitai_image_id
|
from ..utils.civitai_utils import extract_civitai_image_id
|
||||||
from ..utils.utils import calculate_recipe_fingerprint, fuzzy_match
|
from ..utils.utils import calculate_recipe_fingerprint
|
||||||
from natsort import natsorted
|
from natsort import natsorted
|
||||||
import sys
|
import sys
|
||||||
import re
|
import re
|
||||||
@@ -1020,13 +1020,16 @@ class RecipeScanner:
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
result = self._fts_index.search(search, fields)
|
result = self._fts_index.search(search, fields)
|
||||||
# Return None if empty to trigger fuzzy fallback
|
# Return empty set for empty FTS results — do NOT fall back to
|
||||||
# Empty FTS results may indicate query syntax issues or need for fuzzy matching
|
# Python fuzzy matching, which freezes the server with 10k+ recipes.
|
||||||
|
# FTS5 prefix matching with unicode61 tokenizer correctly handles
|
||||||
|
# compound tokens (e.g. "illustrious" matches "path/illustrious/model").
|
||||||
|
# If FTS returns nothing, there are genuinely no matching recipes.
|
||||||
if not result:
|
if not result:
|
||||||
return None
|
return set()
|
||||||
return result
|
return result
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.debug("FTS search failed, falling back to fuzzy search: %s", exc)
|
logger.debug("FTS search failed, falling back to title-only search: %s", exc)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _update_fts_index_for_recipe(
|
def _update_fts_index_for_recipe(
|
||||||
@@ -2079,49 +2082,14 @@ class RecipeScanner:
|
|||||||
if str(item.get("id", "")) in fts_matching_ids
|
if str(item.get("id", "")) in fts_matching_ids
|
||||||
]
|
]
|
||||||
else:
|
else:
|
||||||
# Fallback to fuzzy_match (slower but always available)
|
# FTS index not yet built — return empty rather than
|
||||||
# Build the search predicate based on search options
|
# scanning 42k+ items in Python. The FTS background build
|
||||||
def matches_search(item):
|
# finishes in seconds; by the time a user navigates here
|
||||||
# Search in title if enabled
|
# and types a search, it is already available.
|
||||||
if search_options.get("title", True):
|
logger.debug(
|
||||||
if fuzzy_match(str(item.get("title", "")), search):
|
"FTS index not ready — search '%s' returning empty", search
|
||||||
return True
|
)
|
||||||
|
filtered_data = []
|
||||||
# Search in tags if enabled
|
|
||||||
if search_options.get("tags", True) and "tags" in item:
|
|
||||||
for tag in item["tags"]:
|
|
||||||
if fuzzy_match(tag, search):
|
|
||||||
return True
|
|
||||||
|
|
||||||
# Search in lora file names if enabled
|
|
||||||
if search_options.get("lora_name", True) and "loras" in item:
|
|
||||||
for lora in item["loras"]:
|
|
||||||
if fuzzy_match(str(lora.get("file_name", "")), search):
|
|
||||||
return True
|
|
||||||
|
|
||||||
# Search in lora model names if enabled
|
|
||||||
if search_options.get("lora_model", True) and "loras" in item:
|
|
||||||
for lora in item["loras"]:
|
|
||||||
if fuzzy_match(str(lora.get("modelName", "")), search):
|
|
||||||
return True
|
|
||||||
|
|
||||||
# Search in prompt and negative_prompt if enabled
|
|
||||||
if search_options.get("prompt", True) and "gen_params" in item:
|
|
||||||
gen_params = item["gen_params"]
|
|
||||||
if fuzzy_match(str(gen_params.get("prompt", "")), search):
|
|
||||||
return True
|
|
||||||
if fuzzy_match(
|
|
||||||
str(gen_params.get("negative_prompt", "")), search
|
|
||||||
):
|
|
||||||
return True
|
|
||||||
|
|
||||||
# No match found
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Filter the data using the search predicate
|
|
||||||
filtered_data = [
|
|
||||||
item for item in filtered_data if matches_search(item)
|
|
||||||
]
|
|
||||||
|
|
||||||
# Apply additional filters
|
# Apply additional filters
|
||||||
if filters:
|
if filters:
|
||||||
|
|||||||
@@ -216,11 +216,12 @@ class RecipePersistenceService:
|
|||||||
"preview_nsfw_level",
|
"preview_nsfw_level",
|
||||||
"favorite",
|
"favorite",
|
||||||
"gen_params",
|
"gen_params",
|
||||||
|
"base_model",
|
||||||
)
|
)
|
||||||
|
|
||||||
if not any(key in updates for key in allowed_fields):
|
if not any(key in updates for key in allowed_fields):
|
||||||
raise RecipeValidationError(
|
raise RecipeValidationError(
|
||||||
"At least one field to update must be provided (title or tags or source_path or preview_nsfw_level or favorite or gen_params)"
|
"At least one field to update must be provided (title or tags or source_path or preview_nsfw_level or favorite or gen_params or base_model)"
|
||||||
)
|
)
|
||||||
|
|
||||||
if "gen_params" in updates and not isinstance(updates["gen_params"], dict):
|
if "gen_params" in updates and not isinstance(updates["gen_params"], dict):
|
||||||
|
|||||||
@@ -65,6 +65,8 @@ DEFAULT_SETTINGS: Dict[str, Any] = {
|
|||||||
"onboarding_completed": False,
|
"onboarding_completed": False,
|
||||||
"dismissed_banners": [],
|
"dismissed_banners": [],
|
||||||
"enable_metadata_archive_db": False,
|
"enable_metadata_archive_db": False,
|
||||||
|
"enable_civarchive_api": True,
|
||||||
|
"metadata_provider_order": "civitai_archive_sqlite",
|
||||||
"proxy_enabled": False,
|
"proxy_enabled": False,
|
||||||
"proxy_host": "",
|
"proxy_host": "",
|
||||||
"proxy_port": "",
|
"proxy_port": "",
|
||||||
@@ -107,6 +109,11 @@ DEFAULT_SETTINGS: Dict[str, Any] = {
|
|||||||
"backup_retention_count": 5,
|
"backup_retention_count": 5,
|
||||||
"use_new_license_icons": True,
|
"use_new_license_icons": True,
|
||||||
"group_by_model": False,
|
"group_by_model": False,
|
||||||
|
# AI / LLM provider configuration (BYOK)
|
||||||
|
"llm_provider": "openai", # "openai" | "ollama" | "custom"
|
||||||
|
"llm_api_key": "",
|
||||||
|
"llm_api_base": "", # empty = provider default
|
||||||
|
"llm_model": "", # e.g. "gpt-4o-mini"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -147,6 +154,11 @@ class SettingsManager:
|
|||||||
self._check_environment_variables()
|
self._check_environment_variables()
|
||||||
self._collect_configuration_warnings()
|
self._collect_configuration_warnings()
|
||||||
|
|
||||||
|
if os.environ.get("LORA_MANAGER_PORTABLE", "0") == "1":
|
||||||
|
if not self.settings.get("use_portable_settings"):
|
||||||
|
self.settings["use_portable_settings"] = True
|
||||||
|
self._save_settings()
|
||||||
|
|
||||||
if self._needs_initial_save:
|
if self._needs_initial_save:
|
||||||
self._save_settings()
|
self._save_settings()
|
||||||
self._needs_initial_save = False
|
self._needs_initial_save = False
|
||||||
@@ -620,12 +632,37 @@ class SettingsManager:
|
|||||||
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _normalize_path_set(paths: Iterable[str]) -> set[str]:
|
||||||
|
"""Normalize an iterable of paths for set-based overlap comparison.
|
||||||
|
|
||||||
|
Resolves symlinks via ``os.path.realpath`` when the path exists on disk,
|
||||||
|
then applies ``os.path.normcase`` + ``os.path.normpath`` for consistent
|
||||||
|
cross-platform comparison. Non-string / empty entries are skipped.
|
||||||
|
"""
|
||||||
|
result: set[str] = set()
|
||||||
|
for p in paths:
|
||||||
|
if not isinstance(p, str):
|
||||||
|
continue
|
||||||
|
stripped = p.strip()
|
||||||
|
if not stripped:
|
||||||
|
continue
|
||||||
|
if os.path.exists(stripped):
|
||||||
|
stripped = os.path.normpath(os.path.realpath(stripped))
|
||||||
|
result.add(os.path.normcase(stripped))
|
||||||
|
return result
|
||||||
|
|
||||||
def _validate_folder_paths(
|
def _validate_folder_paths(
|
||||||
self,
|
self,
|
||||||
library_name: str,
|
library_name: str,
|
||||||
folder_paths: Mapping[str, Iterable[str]],
|
folder_paths: Mapping[str, Iterable[str]],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Ensure folder paths do not overlap with other libraries."""
|
"""Ensure folder paths do not overlap with other libraries.
|
||||||
|
|
||||||
|
Also detects checkpoints ↔ unet path overlap within the same library
|
||||||
|
(including via symlink resolution), which is a configuration error since
|
||||||
|
these model types must use separate physical folders.
|
||||||
|
"""
|
||||||
libraries = self.settings.get("libraries", {})
|
libraries = self.settings.get("libraries", {})
|
||||||
normalized_new: Dict[str, Dict[str, str]] = {}
|
normalized_new: Dict[str, Dict[str, str]] = {}
|
||||||
for key, values in folder_paths.items():
|
for key, values in folder_paths.items():
|
||||||
@@ -663,6 +700,22 @@ class SettingsManager:
|
|||||||
f"Folder path(s) {collisions} already assigned to library '{other_name}'"
|
f"Folder path(s) {collisions} already assigned to library '{other_name}'"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Checkpoints ↔ unet overlap within the same library
|
||||||
|
ckpt_paths = folder_paths.get("checkpoints", []) or []
|
||||||
|
unet_paths = folder_paths.get("unet", []) or []
|
||||||
|
if ckpt_paths and unet_paths:
|
||||||
|
ckpt_real = self._normalize_path_set(ckpt_paths)
|
||||||
|
unet_real = self._normalize_path_set(unet_paths)
|
||||||
|
overlap = ckpt_real & unet_real
|
||||||
|
if overlap:
|
||||||
|
collisions = ", ".join(sorted(overlap))
|
||||||
|
raise ValueError(
|
||||||
|
f"Path(s) {collisions} are configured for both "
|
||||||
|
f"'checkpoints' and 'unet' (diffusion models). "
|
||||||
|
f"These model types must use separate physical folders. "
|
||||||
|
f"Please remove one of the conflicting entries."
|
||||||
|
)
|
||||||
|
|
||||||
def _update_active_library_entry(
|
def _update_active_library_entry(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -873,6 +926,23 @@ class SettingsManager:
|
|||||||
self.settings["civitai_api_key"] = env_api_key
|
self.settings["civitai_api_key"] = env_api_key
|
||||||
self._save_settings()
|
self._save_settings()
|
||||||
|
|
||||||
|
# LLM provider overrides
|
||||||
|
llm_env_map = {
|
||||||
|
"LLM_API_KEY": "llm_api_key",
|
||||||
|
"LLM_MODEL": "llm_model",
|
||||||
|
"LLM_API_BASE": "llm_api_base",
|
||||||
|
"LLM_PROVIDER": "llm_provider",
|
||||||
|
}
|
||||||
|
llm_changed = False
|
||||||
|
for env_var, settings_key in llm_env_map.items():
|
||||||
|
env_val = os.environ.get(env_var)
|
||||||
|
if env_val:
|
||||||
|
logger.info("Found %s environment variable", env_var)
|
||||||
|
self.settings[settings_key] = env_val
|
||||||
|
llm_changed = True
|
||||||
|
if llm_changed:
|
||||||
|
self._save_settings()
|
||||||
|
|
||||||
def _default_settings_actions(self) -> List[Dict[str, Any]]:
|
def _default_settings_actions(self) -> List[Dict[str, Any]]:
|
||||||
return [
|
return [
|
||||||
{
|
{
|
||||||
@@ -1520,8 +1590,12 @@ class SettingsManager:
|
|||||||
portable_switch_pending = True
|
portable_switch_pending = True
|
||||||
self._prepare_portable_switch(value)
|
self._prepare_portable_switch(value)
|
||||||
if key == "folder_paths" and isinstance(value, Mapping):
|
if key == "folder_paths" and isinstance(value, Mapping):
|
||||||
|
active_name = self.get_active_library_name()
|
||||||
|
self._validate_folder_paths(active_name, value)
|
||||||
self._update_active_library_entry(folder_paths=value) # type: ignore[arg-type]
|
self._update_active_library_entry(folder_paths=value) # type: ignore[arg-type]
|
||||||
elif key == "extra_folder_paths" and isinstance(value, Mapping):
|
elif key == "extra_folder_paths" and isinstance(value, Mapping):
|
||||||
|
active_name = self.get_active_library_name()
|
||||||
|
self._validate_folder_paths(active_name, value)
|
||||||
self._update_active_library_entry(extra_folder_paths=value) # type: ignore[arg-type]
|
self._update_active_library_entry(extra_folder_paths=value) # type: ignore[arg-type]
|
||||||
elif key == "default_lora_root":
|
elif key == "default_lora_root":
|
||||||
self._update_active_library_entry(default_lora_root=str(value))
|
self._update_active_library_entry(default_lora_root=str(value))
|
||||||
@@ -1775,6 +1849,9 @@ class SettingsManager:
|
|||||||
if key in self.settings:
|
if key in self.settings:
|
||||||
minimal[key] = copy.deepcopy(self.settings[key])
|
minimal[key] = copy.deepcopy(self.settings[key])
|
||||||
|
|
||||||
|
if self.settings.get("use_portable_settings"):
|
||||||
|
minimal["use_portable_settings"] = True
|
||||||
|
|
||||||
if self._seed_template:
|
if self._seed_template:
|
||||||
for key, value in self._seed_template.items():
|
for key, value in self._seed_template.items():
|
||||||
minimal.setdefault(key, copy.deepcopy(value))
|
minimal.setdefault(key, copy.deepcopy(value))
|
||||||
|
|||||||
@@ -51,6 +51,10 @@ class BulkMetadataRefreshUseCase:
|
|||||||
if not model.get("skip_metadata_refresh", False)
|
if not model.get("skip_metadata_refresh", False)
|
||||||
and not self._is_in_skip_path(model.get("folder", ""), skip_paths)
|
and not self._is_in_skip_path(model.get("folder", ""), skip_paths)
|
||||||
and (not model.get("civitai") or not model["civitai"].get("id"))
|
and (not model.get("civitai") or not model["civitai"].get("id"))
|
||||||
|
# Skip models downloaded from Hugging Face — they are not on
|
||||||
|
# CivitAI / CivArchive. Users can still refresh them individually
|
||||||
|
# via the right-click context menu.
|
||||||
|
and not model.get("hf_url", "")
|
||||||
and not (
|
and not (
|
||||||
# Skip models confirmed not on CivitAI when no need to retry
|
# Skip models confirmed not on CivitAI when no need to retry
|
||||||
model.get("from_civitai") is False
|
model.get("from_civitai") is False
|
||||||
@@ -122,6 +126,7 @@ class BulkMetadataRefreshUseCase:
|
|||||||
if sha256:
|
if sha256:
|
||||||
model["sha256"] = sha256
|
model["sha256"] = sha256
|
||||||
model["hash_status"] = "completed"
|
model["hash_status"] = "completed"
|
||||||
|
hash_status = "completed"
|
||||||
else:
|
else:
|
||||||
self._logger.error(f"Failed to calculate hash for {file_path}")
|
self._logger.error(f"Failed to calculate hash for {file_path}")
|
||||||
failures.append({"name": model.get("model_name", file_path or "Unknown"), "error": "Failed to calculate hash"})
|
failures.append({"name": model.get("model_name", file_path or "Unknown"), "error": "Failed to calculate hash"})
|
||||||
@@ -144,6 +149,16 @@ class BulkMetadataRefreshUseCase:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
await MetadataManager.hydrate_model_data(model)
|
await MetadataManager.hydrate_model_data(model)
|
||||||
|
|
||||||
|
# hydrate_model_data replaces model with .metadata.json content,
|
||||||
|
# which may lack sha256. Restore from cache and persist the fix.
|
||||||
|
if not model.get("sha256"):
|
||||||
|
model["sha256"] = sha256
|
||||||
|
model["hash_status"] = model.get("hash_status", hash_status)
|
||||||
|
data_to_save = model.copy()
|
||||||
|
data_to_save.pop("folder", None)
|
||||||
|
await MetadataManager.save_metadata(file_path, data_to_save)
|
||||||
|
|
||||||
result, error_msg = await self._metadata_sync.fetch_and_update_model(
|
result, error_msg = await self._metadata_sync.fetch_and_update_model(
|
||||||
sha256=model["sha256"],
|
sha256=model["sha256"],
|
||||||
file_path=model["file_path"],
|
file_path=model["file_path"],
|
||||||
|
|||||||
@@ -19,7 +19,7 @@ logger = logging.getLogger(__name__)
|
|||||||
_WILDCARD_PATTERN = re.compile(r"__([\w\s.\-+/*\\]+?)__")
|
_WILDCARD_PATTERN = re.compile(r"__([\w\s.\-+/*\\]+?)__")
|
||||||
_OPTION_PATTERN = re.compile(r"{([^{}]*?)}")
|
_OPTION_PATTERN = re.compile(r"{([^{}]*?)}")
|
||||||
_TRIGGER_WORD_PATTERN = re.compile(r"^trigger_words\d+$")
|
_TRIGGER_WORD_PATTERN = re.compile(r"^trigger_words\d+$")
|
||||||
_WEIGHTED_OPTION_PATTERN = re.compile(r"^\s*([0-9.]+)::")
|
_WEIGHTED_OPTION_PATTERN = re.compile(r"^\s*-?\d+(\.\d+)?::")
|
||||||
_NUMERIC_PATTERN = re.compile(r"^-?\d+(\.\d+)?$")
|
_NUMERIC_PATTERN = re.compile(r"^-?\d+(\.\d+)?$")
|
||||||
|
|
||||||
|
|
||||||
@@ -390,7 +390,7 @@ class WildcardService:
|
|||||||
) -> str | None:
|
) -> str | None:
|
||||||
keyword = _normalize_wildcard_key(raw_key)
|
keyword = _normalize_wildcard_key(raw_key)
|
||||||
if keyword in wildcard_dict:
|
if keyword in wildcard_dict:
|
||||||
return rng.choice(wildcard_dict[keyword])
|
return self._pick_weighted_or_plain(wildcard_dict[keyword], rng)
|
||||||
|
|
||||||
if "*" in keyword:
|
if "*" in keyword:
|
||||||
regex_pattern = keyword.replace("*", ".*").replace("+", r"\+")
|
regex_pattern = keyword.replace("*", ".*").replace("+", r"\+")
|
||||||
@@ -400,7 +400,7 @@ class WildcardService:
|
|||||||
if compiled.match(key):
|
if compiled.match(key):
|
||||||
aggregated.extend(values)
|
aggregated.extend(values)
|
||||||
if aggregated:
|
if aggregated:
|
||||||
return rng.choice(aggregated)
|
return self._pick_weighted_or_plain(aggregated, rng)
|
||||||
|
|
||||||
if "/" not in keyword:
|
if "/" not in keyword:
|
||||||
fallback_keyword = _normalize_wildcard_key(f"*/{keyword}")
|
fallback_keyword = _normalize_wildcard_key(f"*/{keyword}")
|
||||||
@@ -409,6 +409,39 @@ class WildcardService:
|
|||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def _pick_weighted_or_plain(
|
||||||
|
self, values: list[str], rng: random.Random
|
||||||
|
) -> str:
|
||||||
|
"""Pick a value from the list, respecting N::weight prefix if present.
|
||||||
|
|
||||||
|
When any value in the list uses the ``N::value`` weighted syntax with a
|
||||||
|
weight different from 1, the pick uses weighted random selection. When
|
||||||
|
no such weighting is present, a plain ``rng.choice`` is used (preserving
|
||||||
|
backward compatibility for unweighted wildcard files).
|
||||||
|
|
||||||
|
In either case the ``N::`` prefix is always stripped from the returned
|
||||||
|
value, matching the behaviour of ``{...}`` option groups.
|
||||||
|
"""
|
||||||
|
# Fast path: skip weighting logic entirely when no :: syntax exists
|
||||||
|
if not any("::" in v for v in values):
|
||||||
|
return rng.choice(values)
|
||||||
|
|
||||||
|
weighted_options: list[tuple[float, str]] = []
|
||||||
|
for value in values:
|
||||||
|
weight = 1.0
|
||||||
|
parts = value.split("::", 1)
|
||||||
|
if len(parts) == 2 and _is_numeric_string(parts[0].strip()):
|
||||||
|
weight = float(parts[0].strip())
|
||||||
|
weighted_options.append((weight, value))
|
||||||
|
|
||||||
|
any_weighted = any(w != 1.0 for w, _ in weighted_options)
|
||||||
|
if any_weighted:
|
||||||
|
picked = self._weighted_choice(weighted_options, rng)
|
||||||
|
else:
|
||||||
|
picked = rng.choice(values)
|
||||||
|
|
||||||
|
return self._strip_weight_prefix(picked)
|
||||||
|
|
||||||
|
|
||||||
def is_trigger_words_input(name: str) -> bool:
|
def is_trigger_words_input(name: str) -> bool:
|
||||||
return bool(_TRIGGER_WORD_PATTERN.match(name))
|
return bool(_TRIGGER_WORD_PATTERN.match(name))
|
||||||
|
|||||||
+14
-1
@@ -12,6 +12,7 @@ NODE_TYPES = {
|
|||||||
"Lora Loader (LoraManager)": 1,
|
"Lora Loader (LoraManager)": 1,
|
||||||
"Lora Stacker (LoraManager)": 2,
|
"Lora Stacker (LoraManager)": 2,
|
||||||
"WanVideo Lora Select (LoraManager)": 3,
|
"WanVideo Lora Select (LoraManager)": 3,
|
||||||
|
"Create Hook LoRA (LoraManager)": 4,
|
||||||
}
|
}
|
||||||
|
|
||||||
# Default ComfyUI node color when bgcolor is null
|
# Default ComfyUI node color when bgcolor is null
|
||||||
@@ -226,9 +227,21 @@ SUPPORTED_DOWNLOAD_SKIP_BASE_MODELS = frozenset(
|
|||||||
"Wan Video 2.5 I2V",
|
"Wan Video 2.5 I2V",
|
||||||
"Hunyuan Video",
|
"Hunyuan Video",
|
||||||
"Anima",
|
"Anima",
|
||||||
|
"ACE Audio",
|
||||||
|
"Boogu",
|
||||||
"Ernie",
|
"Ernie",
|
||||||
"Ernie Turbo",
|
"Ernie Turbo",
|
||||||
"Nucleus",
|
"Grok",
|
||||||
|
"HappyHorse",
|
||||||
|
"HiDream-O1",
|
||||||
|
"Ideogram 4.0",
|
||||||
"Krea 2",
|
"Krea 2",
|
||||||
|
"Lens",
|
||||||
|
"MAI",
|
||||||
|
"Nucleus",
|
||||||
|
"Qwen 2",
|
||||||
|
"Upscaler",
|
||||||
|
"Wan Image 2.7",
|
||||||
|
"Wan Video 2.7",
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -475,13 +475,19 @@ class MetadataUpdater:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
model_folder = get_model_folder(model_hash)
|
model_folder = get_model_folder(model_hash)
|
||||||
if not model_folder:
|
if not model_folder or not os.path.isdir(model_folder):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
civitai = getattr(metadata, "civitai", None)
|
civitai = getattr(metadata, "civitai", None)
|
||||||
if not isinstance(civitai, dict):
|
if not isinstance(civitai, dict):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
# Read the directory listing once so every image entry reuses it.
|
||||||
|
try:
|
||||||
|
dir_entries = os.listdir(model_folder)
|
||||||
|
except OSError:
|
||||||
|
dir_entries = []
|
||||||
|
|
||||||
has_changes = False
|
has_changes = False
|
||||||
|
|
||||||
custom_images = civitai.get("customImages")
|
custom_images = civitai.get("customImages")
|
||||||
@@ -493,24 +499,15 @@ class MetadataUpdater:
|
|||||||
if not img_id:
|
if not img_id:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
if not os.path.isdir(model_folder):
|
prefix = f"custom_{img_id}"
|
||||||
|
found = any(
|
||||||
|
f.startswith(prefix) and os.path.isfile(
|
||||||
|
os.path.join(model_folder, f)
|
||||||
|
)
|
||||||
|
for f in dir_entries
|
||||||
|
)
|
||||||
|
if not found:
|
||||||
stale.append(idx)
|
stale.append(idx)
|
||||||
else:
|
|
||||||
found = False
|
|
||||||
try:
|
|
||||||
prefix = f"custom_{img_id}"
|
|
||||||
for fname in os.listdir(model_folder):
|
|
||||||
if fname.startswith(prefix) and os.path.isfile(
|
|
||||||
os.path.join(model_folder, fname)
|
|
||||||
):
|
|
||||||
found = True
|
|
||||||
break
|
|
||||||
except OSError:
|
|
||||||
stale.append(idx)
|
|
||||||
continue
|
|
||||||
|
|
||||||
if not found:
|
|
||||||
stale.append(idx)
|
|
||||||
|
|
||||||
if stale:
|
if stale:
|
||||||
for idx in reversed(stale):
|
for idx in reversed(stale):
|
||||||
@@ -532,22 +529,9 @@ class MetadataUpdater:
|
|||||||
# is gone.
|
# is gone.
|
||||||
continue
|
continue
|
||||||
|
|
||||||
if not os.path.isdir(model_folder):
|
prefix = f"image_{idx}."
|
||||||
|
if not any(f.startswith(prefix) for f in dir_entries):
|
||||||
stale.append(idx)
|
stale.append(idx)
|
||||||
else:
|
|
||||||
found = False
|
|
||||||
try:
|
|
||||||
prefix = f"image_{idx}."
|
|
||||||
for fname in os.listdir(model_folder):
|
|
||||||
if fname.startswith(prefix):
|
|
||||||
found = True
|
|
||||||
break
|
|
||||||
except OSError:
|
|
||||||
stale.append(idx)
|
|
||||||
continue
|
|
||||||
|
|
||||||
if not found:
|
|
||||||
stale.append(idx)
|
|
||||||
|
|
||||||
if stale:
|
if stale:
|
||||||
for idx in reversed(stale):
|
for idx in reversed(stale):
|
||||||
|
|||||||
@@ -3,9 +3,16 @@ import logging
|
|||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
import json
|
import json
|
||||||
|
import shutil
|
||||||
from ..services.settings_manager import get_settings_manager
|
from ..services.settings_manager import get_settings_manager
|
||||||
from ..services.service_registry import ServiceRegistry
|
from ..services.service_registry import ServiceRegistry
|
||||||
from ..utils.example_images_paths import iter_library_roots
|
from ..utils.example_images_paths import (
|
||||||
|
get_example_images_root,
|
||||||
|
is_hash_folder,
|
||||||
|
iter_library_roots,
|
||||||
|
uses_library_scoped_folders,
|
||||||
|
_library_folder_has_only_hash_dirs,
|
||||||
|
)
|
||||||
from ..utils.metadata_manager import MetadataManager
|
from ..utils.metadata_manager import MetadataManager
|
||||||
from ..utils.example_images_processor import ExampleImagesProcessor
|
from ..utils.example_images_processor import ExampleImagesProcessor
|
||||||
from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
|
from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
|
||||||
@@ -36,6 +43,90 @@ settings = _SettingsProxy()
|
|||||||
class ExampleImagesMigration:
|
class ExampleImagesMigration:
|
||||||
"""Handles migrations for example images naming conventions"""
|
"""Handles migrations for example images naming conventions"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _consolidate_library_folders():
|
||||||
|
"""Move hash folders from library-named subdirectories back to root.
|
||||||
|
|
||||||
|
When a user switches from multi-library mode back to single-library
|
||||||
|
mode, example images previously stored under e.g.
|
||||||
|
``<root>/default/<hash>/`` need to be moved back to
|
||||||
|
``<root>/<hash>/``. Running this once at startup removes the need
|
||||||
|
for ``get_model_folder()`` to perform directory scans on every
|
||||||
|
request.
|
||||||
|
"""
|
||||||
|
if uses_library_scoped_folders():
|
||||||
|
return
|
||||||
|
|
||||||
|
root = get_example_images_root()
|
||||||
|
if not root or not os.path.isdir(root):
|
||||||
|
return
|
||||||
|
|
||||||
|
moved: list[str] = []
|
||||||
|
cleaned: list[str] = []
|
||||||
|
|
||||||
|
try:
|
||||||
|
for entry in os.listdir(root):
|
||||||
|
# Fast regex checks first — no filesystem I/O.
|
||||||
|
if is_hash_folder(entry) or entry == "_deleted":
|
||||||
|
continue
|
||||||
|
|
||||||
|
entry_path = os.path.join(root, entry)
|
||||||
|
if not os.path.isdir(entry_path):
|
||||||
|
continue
|
||||||
|
if not _library_folder_has_only_hash_dirs(entry_path):
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
for hash_entry in os.listdir(entry_path):
|
||||||
|
hash_path = os.path.join(entry_path, hash_entry)
|
||||||
|
if not os.path.isdir(hash_path) or not is_hash_folder(hash_entry):
|
||||||
|
continue
|
||||||
|
target = os.path.join(root, hash_entry)
|
||||||
|
if not os.path.exists(target):
|
||||||
|
try:
|
||||||
|
shutil.move(hash_path, target)
|
||||||
|
moved.append(hash_entry)
|
||||||
|
except (OSError, shutil.Error) as exc:
|
||||||
|
logger.error(
|
||||||
|
"Failed to move '%s' → '%s': %s",
|
||||||
|
hash_path, target, exc,
|
||||||
|
)
|
||||||
|
except OSError as exc:
|
||||||
|
logger.error(
|
||||||
|
"Failed to list library subdirectory '%s': %s",
|
||||||
|
entry_path, exc,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
remaining = os.listdir(entry_path)
|
||||||
|
except OSError:
|
||||||
|
remaining = []
|
||||||
|
if not remaining:
|
||||||
|
try:
|
||||||
|
os.rmdir(entry_path)
|
||||||
|
cleaned.append(entry)
|
||||||
|
except OSError as exc:
|
||||||
|
logger.debug(
|
||||||
|
"Could not remove empty library dir '%s': %s",
|
||||||
|
entry_path, exc,
|
||||||
|
)
|
||||||
|
except OSError as exc:
|
||||||
|
logger.error(
|
||||||
|
"Failed to list example images root during consolidation: %s",
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
|
||||||
|
if moved:
|
||||||
|
logger.info(
|
||||||
|
"Consolidated %d example image folder(s) to root",
|
||||||
|
len(moved),
|
||||||
|
)
|
||||||
|
if cleaned:
|
||||||
|
logger.info(
|
||||||
|
"Removed %d empty library directories",
|
||||||
|
len(cleaned),
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def check_and_run_migrations():
|
async def check_and_run_migrations():
|
||||||
"""Check if migrations are needed and run them in background"""
|
"""Check if migrations are needed and run them in background"""
|
||||||
@@ -44,6 +135,10 @@ class ExampleImagesMigration:
|
|||||||
logger.debug("No example images path configured or path doesn't exist, skipping migrations")
|
logger.debug("No example images path configured or path doesn't exist, skipping migrations")
|
||||||
return
|
return
|
||||||
|
|
||||||
|
# Run library-to-root consolidation once at startup so the hot
|
||||||
|
# path (get_model_folder) stays a pure-path computation.
|
||||||
|
ExampleImagesMigration._consolidate_library_folders()
|
||||||
|
|
||||||
for library_name, library_path in iter_library_roots():
|
for library_name, library_path in iter_library_roots():
|
||||||
if not library_path or not os.path.exists(library_path):
|
if not library_path or not os.path.exists(library_path):
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -83,7 +83,12 @@ def ensure_library_root_exists(library_name: Optional[str] = None) -> str:
|
|||||||
|
|
||||||
|
|
||||||
def get_model_folder(model_hash: str, library_name: Optional[str] = None) -> str:
|
def get_model_folder(model_hash: str, library_name: Optional[str] = None) -> str:
|
||||||
"""Return the folder path for a model's example images."""
|
"""Return the folder path for a model's example images.
|
||||||
|
|
||||||
|
Multi-library ↔ single-library consolidation is handled once at startup by
|
||||||
|
``ExampleImagesMigration._consolidate_library_folders`` — this function is a
|
||||||
|
pure path computation on the hot path (no directory scans).
|
||||||
|
"""
|
||||||
|
|
||||||
if not model_hash:
|
if not model_hash:
|
||||||
return ""
|
return ""
|
||||||
|
|||||||
@@ -35,6 +35,9 @@ class BaseModelMetadata:
|
|||||||
metadata_source: Optional[str] = None # Last provider that supplied metadata
|
metadata_source: Optional[str] = None # Last provider that supplied metadata
|
||||||
last_checked_at: float = 0 # Last checked timestamp
|
last_checked_at: float = 0 # Last checked timestamp
|
||||||
hash_status: str = "completed" # Hash calculation status: pending | calculating | completed | failed
|
hash_status: str = "completed" # Hash calculation status: pending | calculating | completed | failed
|
||||||
|
trainedWords: List[str] = field(
|
||||||
|
default_factory=list
|
||||||
|
) # Trigger words / activation prompts (source-agnostic)
|
||||||
_unknown_fields: Dict[str, Any] = field(
|
_unknown_fields: Dict[str, Any] = field(
|
||||||
default_factory=dict, repr=False, compare=False
|
default_factory=dict, repr=False, compare=False
|
||||||
) # Store unknown fields
|
) # Store unknown fields
|
||||||
@@ -47,6 +50,9 @@ class BaseModelMetadata:
|
|||||||
if self.tags is None:
|
if self.tags is None:
|
||||||
self.tags = []
|
self.tags = []
|
||||||
|
|
||||||
|
if self.trainedWords is None:
|
||||||
|
self.trainedWords = []
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_dict(cls, data: Dict) -> "BaseModelMetadata":
|
def from_dict(cls, data: Dict) -> "BaseModelMetadata":
|
||||||
"""Create instance from dictionary"""
|
"""Create instance from dictionary"""
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ from platformdirs import user_config_dir
|
|||||||
|
|
||||||
|
|
||||||
APP_NAME = "ComfyUI-LoRA-Manager"
|
APP_NAME = "ComfyUI-LoRA-Manager"
|
||||||
|
_LM_PORTABLE_ENV = "LORA_MANAGER_PORTABLE"
|
||||||
_LOGGER = logging.getLogger(__name__)
|
_LOGGER = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@@ -100,7 +101,11 @@ def ensure_settings_file(logger: Optional[logging.Logger] = None) -> str:
|
|||||||
|
|
||||||
|
|
||||||
def _should_use_portable_settings(path: str, logger: logging.Logger) -> bool:
|
def _should_use_portable_settings(path: str, logger: logging.Logger) -> bool:
|
||||||
"""Return ``True`` when the repository settings file enables portable mode."""
|
"""Return ``True`` when the env var forces it or the settings file enables it."""
|
||||||
|
|
||||||
|
if os.environ.get(_LM_PORTABLE_ENV, "0") == "1":
|
||||||
|
logger.debug("Portable mode enabled via %s", _LM_PORTABLE_ENV)
|
||||||
|
return True
|
||||||
|
|
||||||
if not os.path.exists(path):
|
if not os.path.exists(path):
|
||||||
return False
|
return False
|
||||||
|
|||||||
@@ -488,6 +488,12 @@ def calculate_relative_path_for_model(
|
|||||||
if model_type == "embedding":
|
if model_type == "embedding":
|
||||||
formatted_path = formatted_path.replace(" ", "_")
|
formatted_path = formatted_path.replace(" ", "_")
|
||||||
|
|
||||||
|
# Sanitize the resolved path to prevent path traversal
|
||||||
|
formatted_path = formatted_path.lstrip("/")
|
||||||
|
while "//" in formatted_path:
|
||||||
|
formatted_path = formatted_path.replace("//", "/")
|
||||||
|
formatted_path = formatted_path.rstrip("/")
|
||||||
|
|
||||||
return formatted_path
|
return formatted_path
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+1
-1
@@ -1,7 +1,7 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "comfyui-lora-manager"
|
name = "comfyui-lora-manager"
|
||||||
description = "Revolutionize your workflow with the ultimate LoRA companion for ComfyUI!"
|
description = "Revolutionize your workflow with the ultimate LoRA companion for ComfyUI!"
|
||||||
version = "1.1.6"
|
version = "1.2.0"
|
||||||
license = {file = "LICENSE"}
|
license = {file = "LICENSE"}
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"aiohttp",
|
"aiohttp",
|
||||||
|
|||||||
@@ -1,6 +1,10 @@
|
|||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
import json
|
import json
|
||||||
|
# Ensure the script's directory is on sys.path so that py.* imports resolve
|
||||||
|
# regardless of the current working directory (e.g. when launched via
|
||||||
|
# ComfyUI's python_embeded from the ComfyUI root directory).
|
||||||
|
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||||
from py.middleware.cache_middleware import cache_control
|
from py.middleware.cache_middleware import cache_control
|
||||||
from py.middleware.error_middleware import api_json_error
|
from py.middleware.error_middleware import api_json_error
|
||||||
from py.utils.settings_paths import ensure_settings_file
|
from py.utils.settings_paths import ensure_settings_file
|
||||||
|
|||||||
@@ -40,6 +40,12 @@
|
|||||||
margin: 3px 0;
|
margin: 3px 0;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
.context-menu-item.disabled {
|
||||||
|
opacity: 0.4;
|
||||||
|
cursor: not-allowed;
|
||||||
|
pointer-events: none;
|
||||||
|
}
|
||||||
|
|
||||||
.context-menu-item.delete-item {
|
.context-menu-item.delete-item {
|
||||||
color: var(--danger-color);
|
color: var(--danger-color);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -151,6 +151,7 @@ body.modal-open {
|
|||||||
.support-section,
|
.support-section,
|
||||||
.changelog-section,
|
.changelog-section,
|
||||||
.update-info,
|
.update-info,
|
||||||
|
.update-channels,
|
||||||
.info-item,
|
.info-item,
|
||||||
.path-preview {
|
.path-preview {
|
||||||
background: var(--surface-subtle);
|
background: var(--surface-subtle);
|
||||||
|
|||||||
@@ -577,13 +577,14 @@
|
|||||||
border: 1px solid var(--border-color);
|
border: 1px solid var(--border-color);
|
||||||
border-radius: var(--border-radius-sm);
|
border-radius: var(--border-radius-sm);
|
||||||
cursor: pointer;
|
cursor: pointer;
|
||||||
transition: var(--transition-base);
|
transition: var(--transition-base), box-shadow var(--transition-fast), transform var(--transition-fast);
|
||||||
background: var(--bg-color);
|
background: var(--bg-color);
|
||||||
}
|
}
|
||||||
|
|
||||||
.file-option:hover {
|
.file-option:hover {
|
||||||
border-color: var(--lora-accent);
|
border-color: var(--lora-accent);
|
||||||
box-shadow: var(--shadow-sm);
|
box-shadow: var(--shadow-md);
|
||||||
|
transform: translateY(-1px);
|
||||||
}
|
}
|
||||||
|
|
||||||
.file-option.selected {
|
.file-option.selected {
|
||||||
@@ -698,10 +699,25 @@
|
|||||||
color: var(--lora-accent);
|
color: var(--lora-accent);
|
||||||
}
|
}
|
||||||
|
|
||||||
/* Batch Preview List */
|
/* BUG 1 FIX: Single scrollbar — modal-content becomes a flex column so the
|
||||||
|
batch preview step can flex; the list scrolls instead of the modal-content. */
|
||||||
|
#downloadModal .modal-content {
|
||||||
|
display: flex;
|
||||||
|
flex-direction: column;
|
||||||
|
}
|
||||||
|
|
||||||
|
#batchPreviewStep {
|
||||||
|
display: flex;
|
||||||
|
flex-direction: column;
|
||||||
|
min-height: 0;
|
||||||
|
flex: 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
/* Batch Preview List — no max-height; flexes inside #batchPreviewStep */
|
||||||
.batch-preview-list {
|
.batch-preview-list {
|
||||||
max-height: 400px;
|
flex: 1;
|
||||||
overflow-y: auto;
|
overflow-y: auto;
|
||||||
|
min-height: 0;
|
||||||
margin: var(--space-2) 0;
|
margin: var(--space-2) 0;
|
||||||
display: flex;
|
display: flex;
|
||||||
flex-direction: column;
|
flex-direction: column;
|
||||||
@@ -859,6 +875,8 @@
|
|||||||
position: sticky;
|
position: sticky;
|
||||||
top: 0;
|
top: 0;
|
||||||
z-index: 1;
|
z-index: 1;
|
||||||
|
backdrop-filter: blur(8px);
|
||||||
|
-webkit-backdrop-filter: blur(8px);
|
||||||
}
|
}
|
||||||
|
|
||||||
.batch-preview-select-all input[type="checkbox"] {
|
.batch-preview-select-all input[type="checkbox"] {
|
||||||
@@ -884,3 +902,100 @@
|
|||||||
[data-theme="dark"] .batch-preview-select-all {
|
[data-theme="dark"] .batch-preview-select-all {
|
||||||
background: var(--lora-surface);
|
background: var(--lora-surface);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/* FEATURE 2: HF repo grouping — collapsible groups by repo */
|
||||||
|
.batch-preview-group {
|
||||||
|
display: flex;
|
||||||
|
flex-direction: column;
|
||||||
|
background: var(--surface-base);
|
||||||
|
}
|
||||||
|
|
||||||
|
.batch-preview-group-header {
|
||||||
|
display: flex;
|
||||||
|
align-items: center;
|
||||||
|
gap: 8px;
|
||||||
|
padding: 10px 12px;
|
||||||
|
background: var(--color-accent-subtle);
|
||||||
|
border-bottom: 1px solid var(--color-accent-border);
|
||||||
|
cursor: pointer;
|
||||||
|
user-select: none;
|
||||||
|
transition: background var(--transition-fast);
|
||||||
|
}
|
||||||
|
|
||||||
|
.batch-preview-group-header:hover {
|
||||||
|
background: oklch(var(--color-accent-l) var(--color-accent-c) var(--color-accent-h) / 0.18);
|
||||||
|
}
|
||||||
|
|
||||||
|
.batch-preview-group-toggle {
|
||||||
|
width: 14px;
|
||||||
|
font-size: 0.75em;
|
||||||
|
color: var(--text-color);
|
||||||
|
opacity: 0.7;
|
||||||
|
transition: transform var(--transition-fast);
|
||||||
|
flex-shrink: 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
.batch-preview-group-toggle.expanded {
|
||||||
|
transform: rotate(90deg);
|
||||||
|
}
|
||||||
|
|
||||||
|
.batch-preview-group-name {
|
||||||
|
flex: 1;
|
||||||
|
min-width: 0;
|
||||||
|
font-weight: 600;
|
||||||
|
color: var(--text-color);
|
||||||
|
white-space: nowrap;
|
||||||
|
overflow: hidden;
|
||||||
|
text-overflow: ellipsis;
|
||||||
|
font-size: 0.95em;
|
||||||
|
}
|
||||||
|
|
||||||
|
.batch-preview-group-count {
|
||||||
|
font-size: 0.8em;
|
||||||
|
color: var(--text-color);
|
||||||
|
opacity: 0.7;
|
||||||
|
flex-shrink: 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
.batch-preview-group-select-all {
|
||||||
|
width: 18px;
|
||||||
|
height: 18px;
|
||||||
|
cursor: pointer;
|
||||||
|
accent-color: var(--lora-accent);
|
||||||
|
flex-shrink: 0;
|
||||||
|
padding: 0;
|
||||||
|
margin: 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
.batch-preview-group-body {
|
||||||
|
display: flex;
|
||||||
|
flex-direction: column;
|
||||||
|
gap: 1px;
|
||||||
|
background: var(--border-color);
|
||||||
|
overflow: hidden;
|
||||||
|
max-height: 0;
|
||||||
|
opacity: 0;
|
||||||
|
transition: max-height 0.35s ease, opacity 0.2s ease;
|
||||||
|
}
|
||||||
|
|
||||||
|
.batch-preview-group-body.expanded {
|
||||||
|
opacity: 1;
|
||||||
|
max-height: 9999px; /* rest state: content visible; JS inline style overrides during transitions */
|
||||||
|
}
|
||||||
|
|
||||||
|
/* Dark theme overrides for group styles */
|
||||||
|
[data-theme="dark"] .batch-preview-group {
|
||||||
|
background: var(--surface-base);
|
||||||
|
}
|
||||||
|
|
||||||
|
[data-theme="dark"] .batch-preview-group-header {
|
||||||
|
background: var(--color-accent-subtle);
|
||||||
|
}
|
||||||
|
|
||||||
|
[data-theme="dark"] .batch-preview-group-header:hover {
|
||||||
|
background: oklch(var(--color-accent-l) var(--color-accent-c) var(--color-accent-h) / 0.22);
|
||||||
|
}
|
||||||
|
|
||||||
|
[data-theme="dark"] .batch-preview-group-body {
|
||||||
|
background: var(--border-color);
|
||||||
|
}
|
||||||
|
|||||||
@@ -21,18 +21,22 @@
|
|||||||
margin-bottom: 4px;
|
margin-bottom: 4px;
|
||||||
}
|
}
|
||||||
|
|
||||||
.input-group {
|
#relinkCivitaiModal .input-group,
|
||||||
|
#linkHfModal .input-group {
|
||||||
display: flex;
|
display: flex;
|
||||||
flex-direction: column;
|
flex-direction: column;
|
||||||
margin-bottom: var(--space-2);
|
margin-bottom: var(--space-2);
|
||||||
}
|
}
|
||||||
|
|
||||||
.input-group label {
|
#relinkCivitaiModal .input-group label,
|
||||||
|
#linkHfModal .input-group label {
|
||||||
margin-bottom: var(--space-1);
|
margin-bottom: var(--space-1);
|
||||||
font-weight: 500;
|
font-weight: 500;
|
||||||
}
|
}
|
||||||
|
|
||||||
.input-group input {
|
#relinkCivitaiModal .input-group input,
|
||||||
|
#linkHfModal .input-group input {
|
||||||
|
width: auto;
|
||||||
padding: 8px 12px;
|
padding: 8px 12px;
|
||||||
border-radius: var(--border-radius-xs);
|
border-radius: var(--border-radius-xs);
|
||||||
border: 1px solid var(--border-color);
|
border: 1px solid var(--border-color);
|
||||||
|
|||||||
@@ -1562,6 +1562,29 @@ input:checked + .toggle-slider:before {
|
|||||||
box-shadow: 0 0 0 2px rgba(var(--lora-accent-rgb, 79, 70, 229), 0.1);
|
box-shadow: 0 0 0 2px rgba(var(--lora-accent-rgb, 79, 70, 229), 0.1);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
.extra-folder-path-row .path-controls .extra-folder-path-input.has-error {
|
||||||
|
border-color: var(--lora-error);
|
||||||
|
background-color: rgba(220, 53, 69, 0.08);
|
||||||
|
background-color: rgba(from var(--lora-error) r g b / 0.08);
|
||||||
|
}
|
||||||
|
|
||||||
|
.extra-folder-path-row .path-controls .extra-folder-path-input.has-error:focus {
|
||||||
|
box-shadow: 0 0 0 2px rgba(220, 53, 69, 0.15);
|
||||||
|
box-shadow: 0 0 0 2px rgba(from var(--lora-error) r g b / 0.15);
|
||||||
|
}
|
||||||
|
|
||||||
|
.extra-folder-path-error {
|
||||||
|
color: var(--lora-error);
|
||||||
|
font-size: 0.8em;
|
||||||
|
margin-top: 4px;
|
||||||
|
line-height: 1.4;
|
||||||
|
display: none;
|
||||||
|
}
|
||||||
|
|
||||||
|
.extra-folder-path-error.visible {
|
||||||
|
display: block;
|
||||||
|
}
|
||||||
|
|
||||||
.extra-folder-path-row .path-controls .remove-path-btn {
|
.extra-folder-path-row .path-controls .remove-path-btn {
|
||||||
width: 32px;
|
width: 32px;
|
||||||
height: 32px;
|
height: 32px;
|
||||||
@@ -1592,3 +1615,45 @@ input:checked + .toggle-slider:before {
|
|||||||
animation: settings-highlight-pulse 1.5s ease-in-out 3;
|
animation: settings-highlight-pulse 1.5s ease-in-out 3;
|
||||||
border-radius: var(--border-radius-xs);
|
border-radius: var(--border-radius-xs);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/* ---- Combobox panel for AI Provider settings ---- */
|
||||||
|
/* The panel is appended to <body> by Combobox.js and positioned relative to
|
||||||
|
the enhanced <input>. Styles reuse settings-modal CSS variables. */
|
||||||
|
|
||||||
|
.lm-combobox-panel {
|
||||||
|
position: absolute;
|
||||||
|
z-index: 10002;
|
||||||
|
max-height: 240px;
|
||||||
|
overflow-y: auto;
|
||||||
|
background: var(--lora-surface, #2a2a2a);
|
||||||
|
border: 1px solid var(--border-color, rgba(255, 255, 255, 0.12));
|
||||||
|
border-radius: var(--border-radius-xs, 6px);
|
||||||
|
box-shadow: var(--shadow-elevated, 0 6px 18px rgba(0, 0, 0, 0.45));
|
||||||
|
font-size: 0.95em;
|
||||||
|
color: var(--text-color, rgba(226, 232, 240, 0.9));
|
||||||
|
padding: 4px 0;
|
||||||
|
box-sizing: border-box;
|
||||||
|
}
|
||||||
|
|
||||||
|
.lm-combobox-option {
|
||||||
|
padding: 6px 12px;
|
||||||
|
cursor: pointer;
|
||||||
|
user-select: none;
|
||||||
|
white-space: nowrap;
|
||||||
|
overflow: hidden;
|
||||||
|
text-overflow: ellipsis;
|
||||||
|
}
|
||||||
|
|
||||||
|
.lm-combobox-option:hover,
|
||||||
|
.lm-combobox-option.is-active {
|
||||||
|
background: rgba(from var(--lora-accent) r g b / 0.2);
|
||||||
|
color: var(--lora-accent);
|
||||||
|
}
|
||||||
|
|
||||||
|
.lm-combobox-empty {
|
||||||
|
padding: 8px 12px;
|
||||||
|
color: var(--text-color);
|
||||||
|
opacity: 0.45;
|
||||||
|
font-style: italic;
|
||||||
|
user-select: none;
|
||||||
|
}
|
||||||
|
|||||||
@@ -93,15 +93,13 @@
|
|||||||
.update-content {
|
.update-content {
|
||||||
display: flex;
|
display: flex;
|
||||||
flex-direction: column;
|
flex-direction: column;
|
||||||
gap: var(--space-3);
|
gap: var(--space-2);
|
||||||
}
|
}
|
||||||
|
|
||||||
.update-info {
|
.update-info {
|
||||||
display: flex;
|
display: flex;
|
||||||
justify-content: space-between;
|
justify-content: space-between;
|
||||||
align-items: center;
|
align-items: center;
|
||||||
border-radius: var(--border-radius-sm);
|
|
||||||
padding: var(--space-3);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
.update-info .version-info {
|
.update-info .version-info {
|
||||||
@@ -175,7 +173,6 @@
|
|||||||
border: 1px solid var(--lora-border);
|
border: 1px solid var(--lora-border);
|
||||||
border-radius: var(--border-radius-sm);
|
border-radius: var(--border-radius-sm);
|
||||||
padding: var(--space-2);
|
padding: var(--space-2);
|
||||||
margin: var(--space-2) 0;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
[data-theme="dark"] .update-progress {
|
[data-theme="dark"] .update-progress {
|
||||||
@@ -233,11 +230,6 @@
|
|||||||
}
|
}
|
||||||
|
|
||||||
/* Changelog section */
|
/* Changelog section */
|
||||||
.changelog-section {
|
|
||||||
border-radius: var(--border-radius-sm);
|
|
||||||
padding: var(--space-3);
|
|
||||||
}
|
|
||||||
|
|
||||||
.changelog-section h3 {
|
.changelog-section h3 {
|
||||||
margin-top: 0;
|
margin-top: 0;
|
||||||
margin-bottom: var(--space-2);
|
margin-bottom: var(--space-2);
|
||||||
@@ -349,6 +341,131 @@
|
|||||||
text-decoration: underline;
|
text-decoration: underline;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/* Channel Toggle */
|
||||||
|
.update-channels {
|
||||||
|
}
|
||||||
|
|
||||||
|
.channels-label {
|
||||||
|
font-size: 0.9em;
|
||||||
|
color: var(--text-color);
|
||||||
|
opacity: 0.8;
|
||||||
|
margin-bottom: 8px;
|
||||||
|
}
|
||||||
|
|
||||||
|
.channel-toggle {
|
||||||
|
display: flex;
|
||||||
|
gap: 0;
|
||||||
|
background: var(--lora-surface);
|
||||||
|
border-radius: 8px;
|
||||||
|
padding: 3px;
|
||||||
|
width: fit-content;
|
||||||
|
}
|
||||||
|
|
||||||
|
.channel-btn {
|
||||||
|
display: flex;
|
||||||
|
align-items: center;
|
||||||
|
gap: 6px;
|
||||||
|
padding: 8px 20px;
|
||||||
|
border: none;
|
||||||
|
border-radius: 6px;
|
||||||
|
background: transparent;
|
||||||
|
color: var(--text-secondary, #999);
|
||||||
|
cursor: pointer;
|
||||||
|
font-size: 0.9em;
|
||||||
|
font-weight: 500;
|
||||||
|
transition: all 0.2s ease;
|
||||||
|
white-space: nowrap;
|
||||||
|
}
|
||||||
|
|
||||||
|
.channel-btn:hover {
|
||||||
|
color: var(--text-primary, #ddd);
|
||||||
|
background: rgba(255, 255, 255, 0.04);
|
||||||
|
}
|
||||||
|
|
||||||
|
.channel-btn.active {
|
||||||
|
background: var(--lora-accent, #4285F4);
|
||||||
|
color: #fff;
|
||||||
|
box-shadow: 0 1px 3px rgba(0, 0, 0, 0.2);
|
||||||
|
}
|
||||||
|
|
||||||
|
.channel-btn.active i {
|
||||||
|
color: #fff;
|
||||||
|
}
|
||||||
|
|
||||||
|
.channel-btn i {
|
||||||
|
font-size: 0.85em;
|
||||||
|
}
|
||||||
|
|
||||||
|
/* Channel Switch Confirmation Overlay */
|
||||||
|
.channel-switch-overlay {
|
||||||
|
position: fixed;
|
||||||
|
inset: 0;
|
||||||
|
background: rgba(0, 0, 0, 0.6);
|
||||||
|
display: flex;
|
||||||
|
align-items: center;
|
||||||
|
justify-content: center;
|
||||||
|
z-index: 10000;
|
||||||
|
backdrop-filter: blur(2px);
|
||||||
|
}
|
||||||
|
|
||||||
|
.channel-switch-dialog {
|
||||||
|
background: var(--lora-surface);
|
||||||
|
border: 1px solid var(--border-color, rgba(255, 255, 255, 0.1));
|
||||||
|
border-radius: 12px;
|
||||||
|
padding: 28px 32px;
|
||||||
|
max-width: 420px;
|
||||||
|
width: 90%;
|
||||||
|
box-shadow: 0 8px 32px rgba(0, 0, 0, 0.4);
|
||||||
|
}
|
||||||
|
|
||||||
|
.channel-switch-dialog h3 {
|
||||||
|
margin: 0 0 12px;
|
||||||
|
font-size: 1.1em;
|
||||||
|
color: var(--text-primary, #eee);
|
||||||
|
}
|
||||||
|
|
||||||
|
.channel-switch-dialog p {
|
||||||
|
margin: 0 0 24px;
|
||||||
|
font-size: 0.9em;
|
||||||
|
color: var(--text-secondary, #aaa);
|
||||||
|
line-height: 1.6;
|
||||||
|
}
|
||||||
|
|
||||||
|
.channel-switch-actions {
|
||||||
|
display: flex;
|
||||||
|
justify-content: flex-end;
|
||||||
|
gap: 10px;
|
||||||
|
}
|
||||||
|
|
||||||
|
.channel-switch-cancel {
|
||||||
|
padding: 8px 18px;
|
||||||
|
border: 1px solid var(--border-color, rgba(255, 255, 255, 0.1));
|
||||||
|
border-radius: 6px;
|
||||||
|
background: transparent;
|
||||||
|
color: var(--text-secondary, #aaa);
|
||||||
|
cursor: pointer;
|
||||||
|
font-size: 0.9em;
|
||||||
|
}
|
||||||
|
|
||||||
|
.channel-switch-cancel:hover {
|
||||||
|
background: rgba(255, 255, 255, 0.04);
|
||||||
|
}
|
||||||
|
|
||||||
|
.channel-switch-confirm {
|
||||||
|
padding: 8px 18px;
|
||||||
|
border: none;
|
||||||
|
border-radius: 6px;
|
||||||
|
background: var(--lora-accent, #4285F4);
|
||||||
|
color: #fff;
|
||||||
|
cursor: pointer;
|
||||||
|
font-size: 0.9em;
|
||||||
|
font-weight: 500;
|
||||||
|
}
|
||||||
|
|
||||||
|
.channel-switch-confirm:hover {
|
||||||
|
opacity: 0.9;
|
||||||
|
}
|
||||||
|
|
||||||
/* Update preferences section */
|
/* Update preferences section */
|
||||||
.update-preferences {
|
.update-preferences {
|
||||||
border-top: 1px solid var(--lora-border);
|
border-top: 1px solid var(--lora-border);
|
||||||
|
|||||||
@@ -274,6 +274,11 @@
|
|||||||
font-style: italic;
|
font-style: italic;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/* Inline extra tags (selected but not in top-20/appended after API results) */
|
||||||
|
.filter-tag.extra-tag {
|
||||||
|
border-style: dashed;
|
||||||
|
}
|
||||||
|
|
||||||
/* Ensure solid border and full opacity when active or excluded */
|
/* Ensure solid border and full opacity when active or excluded */
|
||||||
.filter-tag.special-tag.active,
|
.filter-tag.special-tag.active,
|
||||||
.filter-tag.special-tag.exclude {
|
.filter-tag.special-tag.exclude {
|
||||||
|
|||||||
@@ -49,10 +49,6 @@ export const MODEL_CONFIG = {
|
|||||||
* @returns {Object} Object containing all API endpoints for the model type
|
* @returns {Object} Object containing all API endpoints for the model type
|
||||||
*/
|
*/
|
||||||
export function getApiEndpoints(modelType) {
|
export function getApiEndpoints(modelType) {
|
||||||
if (!Object.values(MODEL_TYPES).includes(modelType)) {
|
|
||||||
throw new Error(`Invalid model type: ${modelType}`);
|
|
||||||
}
|
|
||||||
|
|
||||||
return {
|
return {
|
||||||
// Base CRUD operations
|
// Base CRUD operations
|
||||||
list: `/api/lm/${modelType}/list`,
|
list: `/api/lm/${modelType}/list`,
|
||||||
@@ -93,6 +89,7 @@ export function getApiEndpoints(modelType) {
|
|||||||
// Query operations
|
// Query operations
|
||||||
scan: `/api/lm/${modelType}/scan`,
|
scan: `/api/lm/${modelType}/scan`,
|
||||||
topTags: `/api/lm/${modelType}/top-tags`,
|
topTags: `/api/lm/${modelType}/top-tags`,
|
||||||
|
searchTags: `/api/lm/${modelType}/search-tags`,
|
||||||
baseModels: `/api/lm/${modelType}/base-models`,
|
baseModels: `/api/lm/${modelType}/base-models`,
|
||||||
roots: `/api/lm/${modelType}/roots`,
|
roots: `/api/lm/${modelType}/roots`,
|
||||||
folders: `/api/lm/${modelType}/folders`,
|
folders: `/api/lm/${modelType}/folders`,
|
||||||
|
|||||||
@@ -112,6 +112,18 @@ export class BaseModelApiClient {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async cancelDownload(downloadId) {
|
||||||
|
try {
|
||||||
|
const response = await fetch(
|
||||||
|
`${DOWNLOAD_ENDPOINTS.cancelGet}?download_id=${encodeURIComponent(downloadId)}`
|
||||||
|
);
|
||||||
|
return await response.json();
|
||||||
|
} catch (error) {
|
||||||
|
console.error('Error cancelling download:', error);
|
||||||
|
return { success: false, error: error.message };
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
async loadMoreWithVirtualScroll(resetPage = false, updateFolders = false) {
|
async loadMoreWithVirtualScroll(resetPage = false, updateFolders = false) {
|
||||||
const pageState = this.getPageState();
|
const pageState = this.getPageState();
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,394 @@
|
|||||||
|
// Combobox.js — Reusable dropdown-suggestion + free-text input component.
|
||||||
|
//
|
||||||
|
// Enhances an existing <input> element with a dropdown panel that merges static
|
||||||
|
// `presets` with asynchronously fetched options (`fetchOptions`). The input
|
||||||
|
// remains a free-text field — selecting a dropdown option is optional, the
|
||||||
|
// user can always type an arbitrary value.
|
||||||
|
//
|
||||||
|
// Zero dependencies: pure DOM manipulation. Exported on `window.Combobox`
|
||||||
|
// so non-module callers can instantiate it, and as a named ES module export
|
||||||
|
// for callers that import it directly.
|
||||||
|
//
|
||||||
|
// Usage:
|
||||||
|
// const box = new Combobox(inputEl, {
|
||||||
|
// presets: ['masterpiece', 'best quality'],
|
||||||
|
// fetchOptions: async (q) => await fetchSuggestions(q),
|
||||||
|
// placeholder: 'Type a value…',
|
||||||
|
// onSelect: (value) => console.log('chose', value),
|
||||||
|
// });
|
||||||
|
// box.updatePresets(['new', 'presets']);
|
||||||
|
// box.setValue('masterpiece');
|
||||||
|
|
||||||
|
const DEBOUNCE_MS = 300;
|
||||||
|
|
||||||
|
export class Combobox {
|
||||||
|
/**
|
||||||
|
* @param {HTMLInputElement} inputElement Existing <input> to enhance.
|
||||||
|
* @param {Object} options
|
||||||
|
* @param {string[]} [options.presets=[]] Static preset values shown in dropdown.
|
||||||
|
* @param {(inputValue: string) => Promise<string[]>} [options.fetchOptions]
|
||||||
|
* Async function returning dynamic suggestions for the current input.
|
||||||
|
* @param {string} [options.placeholder] Placeholder text for the empty state.
|
||||||
|
* @param {(value: string) => void} [options.onSelect] Callback when an option is chosen.
|
||||||
|
*/
|
||||||
|
constructor(inputElement, options = {}) {
|
||||||
|
if (!inputElement || inputElement.tagName !== 'INPUT') {
|
||||||
|
console.error('Combobox: expected an <input> element');
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
this.input = inputElement;
|
||||||
|
this.presets = Array.isArray(options.presets) ? [...options.presets] : [];
|
||||||
|
this.fetchOptions = typeof options.fetchOptions === 'function' ? options.fetchOptions : null;
|
||||||
|
this.placeholder = options.placeholder || '';
|
||||||
|
this.onSelect = typeof options.onSelect === 'function' ? options.onSelect : null;
|
||||||
|
|
||||||
|
// Internal state
|
||||||
|
this._isOpen = false;
|
||||||
|
this._activeIndex = -1;
|
||||||
|
this._renderedOptions = []; // current visible option strings (de-duplicated, merged)
|
||||||
|
this._fetchToken = 0; // guards against out-of-order async fetch results
|
||||||
|
this._fetchTimer = null;
|
||||||
|
this._suppressInputOpen = false; // guards setValue() from reopening the dropdown
|
||||||
|
|
||||||
|
this._buildDropdown();
|
||||||
|
this._bindEvents();
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- public API ----
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Replace the preset list. Re-renders the dropdown if it is open.
|
||||||
|
* @param {string[]} presets
|
||||||
|
* @returns {void}
|
||||||
|
*/
|
||||||
|
updatePresets(presets) {
|
||||||
|
this.presets = Array.isArray(presets) ? [...presets] : [];
|
||||||
|
if (this._isOpen) {
|
||||||
|
this._refresh();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Set the input value programmatically without triggering the dropdown
|
||||||
|
* or firing synthetic events.
|
||||||
|
* @param {string} value
|
||||||
|
* @returns {void}
|
||||||
|
*/
|
||||||
|
setValue(value) {
|
||||||
|
const prev = this._suppressInputOpen;
|
||||||
|
this._suppressInputOpen = true;
|
||||||
|
this.input.value = value ?? '';
|
||||||
|
this._suppressInputOpen = prev;
|
||||||
|
if (this._isOpen) {
|
||||||
|
this._refresh();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- build ----
|
||||||
|
|
||||||
|
_buildDropdown() {
|
||||||
|
const panel = document.createElement('div');
|
||||||
|
panel.className = 'lm-combobox-panel';
|
||||||
|
panel.setAttribute('role', 'listbox');
|
||||||
|
panel.style.display = 'none';
|
||||||
|
// Append to <body> so the panel is never clipped by an overflow:hidden
|
||||||
|
// ancestor; positioning is recomputed on each open.
|
||||||
|
document.body.appendChild(panel);
|
||||||
|
this.panel = panel;
|
||||||
|
|
||||||
|
if (this.placeholder) {
|
||||||
|
this.input.setAttribute('placeholder', this.placeholder);
|
||||||
|
}
|
||||||
|
this.input.setAttribute('autocomplete', 'off');
|
||||||
|
this.input.setAttribute('role', 'combobox');
|
||||||
|
this.input.setAttribute('aria-autocomplete', 'list');
|
||||||
|
this.input.setAttribute('aria-expanded', 'false');
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- event wiring ----
|
||||||
|
|
||||||
|
_bindEvents() {
|
||||||
|
this.input.addEventListener('focus', () => {
|
||||||
|
if (this._suppressInputOpen) return;
|
||||||
|
this._open();
|
||||||
|
});
|
||||||
|
|
||||||
|
this.input.addEventListener('input', () => {
|
||||||
|
if (this._suppressInputOpen) return;
|
||||||
|
this._open(); // no-op if already open
|
||||||
|
this._refresh(); // re-filter by current input value
|
||||||
|
this._scheduleFetch();
|
||||||
|
});
|
||||||
|
|
||||||
|
this.input.addEventListener('keydown', (event) => this._onKeyDown(event));
|
||||||
|
|
||||||
|
// Click an option (delegated)
|
||||||
|
this.panel.addEventListener('click', (event) => {
|
||||||
|
const item = event.target.closest('.lm-combobox-option');
|
||||||
|
if (!item) return;
|
||||||
|
const value = item.dataset.value;
|
||||||
|
if (value !== undefined) {
|
||||||
|
this._choose(value);
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
// Hover updates the active highlight so keyboard + mouse stay in sync.
|
||||||
|
this.panel.addEventListener('mouseover', (event) => {
|
||||||
|
const item = event.target.closest('.lm-combobox-option');
|
||||||
|
if (!item) return;
|
||||||
|
const idx = Number(item.dataset.index);
|
||||||
|
if (!Number.isNaN(idx)) {
|
||||||
|
this._setActiveIndex(idx);
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
// Click outside closes the dropdown.
|
||||||
|
this._outsideClickHandler = (event) => {
|
||||||
|
if (this._isOpen && !this.input.contains(event.target) && !this.panel.contains(event.target)) {
|
||||||
|
this._close();
|
||||||
|
}
|
||||||
|
};
|
||||||
|
document.addEventListener('mousedown', this._outsideClickHandler);
|
||||||
|
|
||||||
|
// Reposition on viewport changes while open.
|
||||||
|
this._resizeHandler = () => {
|
||||||
|
if (this._isOpen) this._position();
|
||||||
|
};
|
||||||
|
window.addEventListener('resize', this._resizeHandler);
|
||||||
|
window.addEventListener('scroll', this._resizeHandler, true);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- keyboard ----
|
||||||
|
|
||||||
|
_onKeyDown(event) {
|
||||||
|
if (!this._isOpen) {
|
||||||
|
if (event.key === 'ArrowDown') {
|
||||||
|
event.preventDefault();
|
||||||
|
this._open();
|
||||||
|
this._setActiveIndex(0);
|
||||||
|
}
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
switch (event.key) {
|
||||||
|
case 'ArrowDown':
|
||||||
|
event.preventDefault();
|
||||||
|
this._setActiveIndex(this._activeIndex + 1);
|
||||||
|
break;
|
||||||
|
|
||||||
|
case 'ArrowUp':
|
||||||
|
event.preventDefault();
|
||||||
|
this._setActiveIndex(this._activeIndex - 1);
|
||||||
|
break;
|
||||||
|
|
||||||
|
case 'Enter':
|
||||||
|
// Only intercept Enter to pick an option when one is actively
|
||||||
|
// highlighted; otherwise let the input's default behavior
|
||||||
|
// (form submit / free-text commit) proceed.
|
||||||
|
if (this._activeIndex >= 0 && this._activeIndex < this._renderedOptions.length) {
|
||||||
|
event.preventDefault();
|
||||||
|
this._choose(this._renderedOptions[this._activeIndex]);
|
||||||
|
}
|
||||||
|
break;
|
||||||
|
|
||||||
|
case 'Escape':
|
||||||
|
event.preventDefault();
|
||||||
|
this._close();
|
||||||
|
this.input.focus();
|
||||||
|
break;
|
||||||
|
|
||||||
|
case 'Tab':
|
||||||
|
// Allow normal tab navigation; just close the panel.
|
||||||
|
this._close();
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- open / close ----
|
||||||
|
|
||||||
|
_open() {
|
||||||
|
if (this._isOpen) return;
|
||||||
|
this._isOpen = true;
|
||||||
|
this.panel.style.display = 'block';
|
||||||
|
this.input.setAttribute('aria-expanded', 'true');
|
||||||
|
// On open, render ALL presets — do not filter by the current input
|
||||||
|
// value. Filtering on the input event is handled separately.
|
||||||
|
this._render(this.presets);
|
||||||
|
this._position();
|
||||||
|
}
|
||||||
|
|
||||||
|
_close() {
|
||||||
|
if (!this._isOpen) return;
|
||||||
|
this._isOpen = false;
|
||||||
|
this.panel.style.display = 'none';
|
||||||
|
this.input.setAttribute('aria-expanded', 'false');
|
||||||
|
this._activeIndex = -1;
|
||||||
|
this._cancelFetch();
|
||||||
|
}
|
||||||
|
|
||||||
|
_position() {
|
||||||
|
const rect = this.input.getBoundingClientRect();
|
||||||
|
const panelHeight = this.panel.offsetHeight;
|
||||||
|
const viewportHeight = window.innerHeight;
|
||||||
|
const spaceBelow = viewportHeight - rect.bottom;
|
||||||
|
const spaceAbove = rect.top;
|
||||||
|
|
||||||
|
// Flip above the input when there is more room there.
|
||||||
|
const placeAbove = spaceBelow < panelHeight && spaceAbove > spaceBelow;
|
||||||
|
const top = placeAbove
|
||||||
|
? rect.top + window.scrollY - panelHeight
|
||||||
|
: rect.bottom + window.scrollY;
|
||||||
|
|
||||||
|
this.panel.style.top = `${Math.max(0, top)}px`;
|
||||||
|
this.panel.style.left = `${rect.left + window.scrollX}px`;
|
||||||
|
this.panel.style.minWidth = `${rect.width}px`;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- rendering ----
|
||||||
|
|
||||||
|
/** Render a list of strings into the panel. */
|
||||||
|
_render(items) {
|
||||||
|
this._renderedOptions = items;
|
||||||
|
this.panel.innerHTML = '';
|
||||||
|
if (items.length === 0) {
|
||||||
|
const empty = document.createElement('div');
|
||||||
|
empty.className = 'lm-combobox-empty';
|
||||||
|
empty.textContent = this.placeholder ? this.placeholder : 'No options';
|
||||||
|
this.panel.appendChild(empty);
|
||||||
|
this._activeIndex = -1;
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
const fragment = document.createDocumentFragment();
|
||||||
|
items.forEach((opt, idx) => {
|
||||||
|
const item = document.createElement('div');
|
||||||
|
item.className = 'lm-combobox-option';
|
||||||
|
item.setAttribute('role', 'option');
|
||||||
|
item.dataset.value = opt;
|
||||||
|
item.dataset.index = String(idx);
|
||||||
|
item.textContent = opt;
|
||||||
|
if (idx === this._activeIndex) {
|
||||||
|
item.classList.add('is-active');
|
||||||
|
}
|
||||||
|
fragment.appendChild(item);
|
||||||
|
});
|
||||||
|
this.panel.appendChild(fragment);
|
||||||
|
|
||||||
|
if (this._activeIndex >= items.length) {
|
||||||
|
this._setActiveIndex(items.length - 1);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Filter presets by current input value and re-render. */
|
||||||
|
_refresh() {
|
||||||
|
const value = this.input.value;
|
||||||
|
const filtered = this._filterPresets(value);
|
||||||
|
const merged = this._mergeUnique(filtered, this._fetchedOptions || []);
|
||||||
|
this._render(merged);
|
||||||
|
}
|
||||||
|
|
||||||
|
_filterPresets(value) {
|
||||||
|
const v = (value || '').toLowerCase();
|
||||||
|
if (!v) return [...this.presets];
|
||||||
|
return this.presets.filter((p) => String(p).toLowerCase().startsWith(v));
|
||||||
|
}
|
||||||
|
|
||||||
|
_mergeUnique(...lists) {
|
||||||
|
const seen = new Set();
|
||||||
|
const out = [];
|
||||||
|
for (const list of lists) {
|
||||||
|
for (const item of list) {
|
||||||
|
const key = String(item);
|
||||||
|
if (!seen.has(key)) {
|
||||||
|
seen.add(key);
|
||||||
|
out.push(key);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out;
|
||||||
|
}
|
||||||
|
|
||||||
|
_setActiveIndex(idx) {
|
||||||
|
const max = this._renderedOptions.length - 1;
|
||||||
|
const clamped = Math.max(-1, Math.min(max, idx));
|
||||||
|
this._activeIndex = clamped;
|
||||||
|
// Update DOM classes without full re-render.
|
||||||
|
const items = this.panel.querySelectorAll('.lm-combobox-option');
|
||||||
|
items.forEach((el, i) => {
|
||||||
|
el.classList.toggle('is-active', i === clamped);
|
||||||
|
});
|
||||||
|
// Scroll the active item into view inside the panel.
|
||||||
|
if (clamped >= 0 && items[clamped]) {
|
||||||
|
items[clamped].scrollIntoView({ block: 'nearest' });
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Remove the panel from the DOM and detach event listeners.
|
||||||
|
* Call this before discarding the Combobox instance.
|
||||||
|
*/
|
||||||
|
destroy() {
|
||||||
|
this._close();
|
||||||
|
if (this.panel && this.panel.parentNode) {
|
||||||
|
this.panel.parentNode.removeChild(this.panel);
|
||||||
|
}
|
||||||
|
document.removeEventListener('mousedown', this._outsideClickHandler);
|
||||||
|
window.removeEventListener('resize', this._resizeHandler);
|
||||||
|
window.removeEventListener('scroll', this._resizeHandler, true);
|
||||||
|
}
|
||||||
|
|
||||||
|
_choose(value) {
|
||||||
|
this.input.value = value;
|
||||||
|
this._close();
|
||||||
|
if (typeof this.onSelect === 'function') {
|
||||||
|
this.onSelect(value);
|
||||||
|
}
|
||||||
|
// Re-focus without reopening the dropdown.
|
||||||
|
this._suppressInputOpen = true;
|
||||||
|
this.input.focus();
|
||||||
|
this._suppressInputOpen = false;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- async fetch (debounced) ----
|
||||||
|
|
||||||
|
_scheduleFetch() {
|
||||||
|
if (!this.fetchOptions) return;
|
||||||
|
this._cancelFetch();
|
||||||
|
this._fetchTimer = setTimeout(() => {
|
||||||
|
this._fetchTimer = null;
|
||||||
|
this._runFetch();
|
||||||
|
}, DEBOUNCE_MS);
|
||||||
|
}
|
||||||
|
|
||||||
|
_cancelFetch() {
|
||||||
|
if (this._fetchTimer) {
|
||||||
|
clearTimeout(this._fetchTimer);
|
||||||
|
this._fetchTimer = null;
|
||||||
|
}
|
||||||
|
this._fetchToken++; // invalidate any in-flight result
|
||||||
|
}
|
||||||
|
|
||||||
|
async _runFetch() {
|
||||||
|
if (!this.fetchOptions) return;
|
||||||
|
const token = this._fetchToken;
|
||||||
|
const value = this.input.value;
|
||||||
|
let results;
|
||||||
|
try {
|
||||||
|
results = await this.fetchOptions(value);
|
||||||
|
} catch (err) {
|
||||||
|
console.error('Combobox fetchOptions error:', err);
|
||||||
|
results = [];
|
||||||
|
}
|
||||||
|
// Stale guard: a newer fetch or close superseded this one.
|
||||||
|
if (token !== this._fetchToken || !this._isOpen) return;
|
||||||
|
this._fetchedOptions = Array.isArray(results) ? results : [];
|
||||||
|
this._refresh();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Expose for non-module callers (templates load via <script type="module">,
|
||||||
|
// but some widget code reads globals off `window`).
|
||||||
|
if (typeof window !== 'undefined') {
|
||||||
|
window.Combobox = Combobox;
|
||||||
|
}
|
||||||
@@ -27,8 +27,9 @@ export class BaseContextMenu {
|
|||||||
const menuItem = e.target.closest('.context-menu-item');
|
const menuItem = e.target.closest('.context-menu-item');
|
||||||
if (!menuItem || !this.currentCard) return;
|
if (!menuItem || !this.currentCard) return;
|
||||||
|
|
||||||
// Ignore clicks on submenu trigger (has-submenu parent)
|
// Ignore clicks on submenu trigger (has-submenu parent) or disabled items
|
||||||
if (menuItem.classList.contains('has-submenu')) return;
|
if (menuItem.classList.contains('has-submenu')) return;
|
||||||
|
if (menuItem.classList.contains('disabled')) return;
|
||||||
|
|
||||||
const action = menuItem.dataset.action;
|
const action = menuItem.dataset.action;
|
||||||
if (!action) return;
|
if (!action) return;
|
||||||
|
|||||||
@@ -274,6 +274,9 @@ export class BulkContextMenu extends BaseContextMenu {
|
|||||||
case 'resume-metadata-refresh':
|
case 'resume-metadata-refresh':
|
||||||
bulkManager.setSkipMetadataRefresh(false);
|
bulkManager.setSkipMetadataRefresh(false);
|
||||||
break;
|
break;
|
||||||
|
case 'enrich-hf-llm-bulk':
|
||||||
|
this.enrichBulkWithAgent();
|
||||||
|
break;
|
||||||
case 'delete-all':
|
case 'delete-all':
|
||||||
bulkManager.showBulkDeleteModal();
|
bulkManager.showBulkDeleteModal();
|
||||||
break;
|
break;
|
||||||
@@ -363,4 +366,90 @@ export class BulkContextMenu extends BaseContextMenu {
|
|||||||
console.error('Bulk download example images failed:', error);
|
console.error('Bulk download example images failed:', error);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Enrich metadata for selected models via LLM agent skill.
|
||||||
|
*/
|
||||||
|
async enrichBulkWithAgent() {
|
||||||
|
if (state.selectedModels.size === 0) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
const { agentManager } = await import('../../managers/AgentManager.js');
|
||||||
|
|
||||||
|
const configured = await agentManager.isLlmConfigured();
|
||||||
|
if (!configured) {
|
||||||
|
showToast('toast.agent.llmNotConfigured', {}, 'warning');
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
const modelPaths = [...state.selectedModels];
|
||||||
|
|
||||||
|
agentManager.connect();
|
||||||
|
|
||||||
|
const progressUI = state.loadingManager.showEnhancedProgress(
|
||||||
|
`Enriching metadata for ${modelPaths.length} models...`
|
||||||
|
);
|
||||||
|
|
||||||
|
function cleanupCallbacks() {
|
||||||
|
const pIdx = agentManager.progressCallbacks.indexOf(onProgress);
|
||||||
|
if (pIdx >= 0) agentManager.progressCallbacks.splice(pIdx, 1);
|
||||||
|
const cIdx = agentManager.completeCallbacks.indexOf(onComplete);
|
||||||
|
if (cIdx >= 0) agentManager.completeCallbacks.splice(cIdx, 1);
|
||||||
|
const eIdx = agentManager.errorCallbacks.indexOf(onError);
|
||||||
|
if (eIdx >= 0) agentManager.errorCallbacks.splice(eIdx, 1);
|
||||||
|
}
|
||||||
|
|
||||||
|
const onProgress = (data) => {
|
||||||
|
if (data.status === 'processing' && data.current_path && data.updated_data && Object.keys(data.updated_data).length > 0) {
|
||||||
|
if (state.virtualScroller?.updateSingleItem) {
|
||||||
|
state.virtualScroller.updateSingleItem(data.current_path, data.updated_data);
|
||||||
|
}
|
||||||
|
const pct = data.total > 0 ? Math.floor((data.processed / data.total) * 100) : 0;
|
||||||
|
const name = data.current_path.split('/').pop();
|
||||||
|
progressUI.updateProgress(pct, name, `Processing ${data.processed}/${data.total}: ${name}`);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
agentManager.onProgress(onProgress);
|
||||||
|
|
||||||
|
const onComplete = (data) => {
|
||||||
|
cleanupCallbacks();
|
||||||
|
|
||||||
|
if (data.status === 'completed') {
|
||||||
|
if (state.bulkMode) bulkManager.toggleBulkMode();
|
||||||
|
progressUI.complete(data.summary || 'Enrich complete');
|
||||||
|
showToast(
|
||||||
|
'toast.agent.enrichComplete',
|
||||||
|
{ summary: data.summary || 'Done' },
|
||||||
|
'success'
|
||||||
|
);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
agentManager.onComplete(onComplete);
|
||||||
|
|
||||||
|
const onError = (data) => {
|
||||||
|
cleanupCallbacks();
|
||||||
|
if (state.bulkMode) bulkManager.toggleBulkMode();
|
||||||
|
state.loadingManager.hide();
|
||||||
|
showToast(
|
||||||
|
'toast.agent.enrichFailed',
|
||||||
|
{ error: data.error || 'Unknown error' },
|
||||||
|
'error'
|
||||||
|
);
|
||||||
|
};
|
||||||
|
agentManager.onError(onError);
|
||||||
|
|
||||||
|
try {
|
||||||
|
await agentManager.executeSkill('enrich_hf_metadata', modelPaths);
|
||||||
|
} catch (error) {
|
||||||
|
cleanupCallbacks();
|
||||||
|
if (state.bulkMode) bulkManager.toggleBulkMode();
|
||||||
|
state.loadingManager.hide();
|
||||||
|
showToast(
|
||||||
|
'toast.agent.enrichFailed',
|
||||||
|
{ error: error.message },
|
||||||
|
'error'
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
import { BaseContextMenu } from './BaseContextMenu.js';
|
import { BaseContextMenu } from './BaseContextMenu.js';
|
||||||
import { ModelContextMenuMixin } from './ModelContextMenuMixin.js';
|
import { ModelContextMenuMixin } from './ModelContextMenuMixin.js';
|
||||||
|
import { state } from '../../state/index.js';
|
||||||
import { getModelApiClient, resetAndReload } from '../../api/modelApiFactory.js';
|
import { getModelApiClient, resetAndReload } from '../../api/modelApiFactory.js';
|
||||||
import { copyLoraSyntax, sendLoraToWorkflow, buildLoraSyntax } from '../../utils/uiHelpers.js';
|
import { copyLoraSyntax, sendLoraToWorkflow, buildLoraSyntax, showToast } from '../../utils/uiHelpers.js';
|
||||||
import { showExcludeModal, showDeleteModal } from '../../utils/modalUtils.js';
|
import { showExcludeModal, showDeleteModal } from '../../utils/modalUtils.js';
|
||||||
import { moveManager } from '../../managers/MoveManager.js';
|
import { moveManager } from '../../managers/MoveManager.js';
|
||||||
|
|
||||||
@@ -23,6 +24,17 @@ export class LoraContextMenu extends BaseContextMenu {
|
|||||||
showMenu(x, y, card) {
|
showMenu(x, y, card) {
|
||||||
super.showMenu(x, y, card);
|
super.showMenu(x, y, card);
|
||||||
this.updateExcludeMenuItem();
|
this.updateExcludeMenuItem();
|
||||||
|
this.updateEnrichMenuItem(card);
|
||||||
|
}
|
||||||
|
|
||||||
|
updateEnrichMenuItem(card) {
|
||||||
|
const enrichItem = this.menu?.querySelector('[data-action="enrich-hf-llm"]');
|
||||||
|
if (!enrichItem) return;
|
||||||
|
const hasHfUrl = !!card.dataset.hf_url;
|
||||||
|
enrichItem.classList.toggle('disabled', !hasHfUrl);
|
||||||
|
enrichItem.title = hasHfUrl
|
||||||
|
? ''
|
||||||
|
: 'Link this model to a HuggingFace repo first (Link Model \u2192 Link to HuggingFace)';
|
||||||
}
|
}
|
||||||
|
|
||||||
handleMenuAction(action, menuItem) {
|
handleMenuAction(action, menuItem) {
|
||||||
@@ -63,6 +75,9 @@ export class LoraContextMenu extends BaseContextMenu {
|
|||||||
case 'refresh-metadata':
|
case 'refresh-metadata':
|
||||||
getModelApiClient().refreshSingleModelMetadata(this.currentCard.dataset.filepath);
|
getModelApiClient().refreshSingleModelMetadata(this.currentCard.dataset.filepath);
|
||||||
break;
|
break;
|
||||||
|
case 'enrich-hf-llm':
|
||||||
|
this.enrichWithAgent(this.currentCard.dataset.filepath);
|
||||||
|
break;
|
||||||
case 'exclude':
|
case 'exclude':
|
||||||
showExcludeModal(this.currentCard.dataset.filepath);
|
showExcludeModal(this.currentCard.dataset.filepath);
|
||||||
break;
|
break;
|
||||||
@@ -72,10 +87,74 @@ export class LoraContextMenu extends BaseContextMenu {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async enrichWithAgent(filePath) {
|
||||||
|
const { agentManager } = await import('../../managers/AgentManager.js');
|
||||||
|
|
||||||
|
const configured = await agentManager.isLlmConfigured();
|
||||||
|
if (!configured) {
|
||||||
|
showToast('toast.agent.llmNotConfigured', {}, 'warning');
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
agentManager.connect();
|
||||||
|
|
||||||
|
const progressUI = state.loadingManager.showEnhancedProgress(
|
||||||
|
'Enriching metadata with AI...'
|
||||||
|
);
|
||||||
|
|
||||||
|
function cleanupCallbacks() {
|
||||||
|
const pIdx = agentManager.progressCallbacks.indexOf(onProgress);
|
||||||
|
if (pIdx >= 0) agentManager.progressCallbacks.splice(pIdx, 1);
|
||||||
|
const cIdx = agentManager.completeCallbacks.indexOf(onComplete);
|
||||||
|
if (cIdx >= 0) agentManager.completeCallbacks.splice(cIdx, 1);
|
||||||
|
const eIdx = agentManager.errorCallbacks.indexOf(onError);
|
||||||
|
if (eIdx >= 0) agentManager.errorCallbacks.splice(eIdx, 1);
|
||||||
|
}
|
||||||
|
|
||||||
|
const onProgress = (data) => {
|
||||||
|
if (data.status === 'processing' && data.current_path && data.updated_data && Object.keys(data.updated_data).length > 0) {
|
||||||
|
if (state.virtualScroller?.updateSingleItem) {
|
||||||
|
state.virtualScroller.updateSingleItem(data.current_path, data.updated_data);
|
||||||
|
}
|
||||||
|
const pct = data.total > 0 ? Math.floor((data.processed / data.total) * 100) : 0;
|
||||||
|
const name = data.current_path.split('/').pop();
|
||||||
|
progressUI.updateProgress(pct, name, `Processing ${name}`);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
agentManager.onProgress(onProgress);
|
||||||
|
|
||||||
|
const onComplete = (data) => {
|
||||||
|
cleanupCallbacks();
|
||||||
|
|
||||||
|
if (data.status === 'completed') {
|
||||||
|
progressUI.complete(data.summary || 'Enrich complete');
|
||||||
|
showToast('toast.agent.enrichComplete', { summary: data.summary || 'Done' }, 'success');
|
||||||
|
}
|
||||||
|
};
|
||||||
|
agentManager.onComplete(onComplete);
|
||||||
|
|
||||||
|
const onError = (data) => {
|
||||||
|
cleanupCallbacks();
|
||||||
|
state.loadingManager.hide();
|
||||||
|
showToast('toast.agent.enrichFailed', { error: data.error || 'Unknown error' }, 'error');
|
||||||
|
};
|
||||||
|
agentManager.onError(onError);
|
||||||
|
|
||||||
|
try {
|
||||||
|
await agentManager.executeSkill('enrich_hf_metadata', [filePath]);
|
||||||
|
} catch (error) {
|
||||||
|
cleanupCallbacks();
|
||||||
|
state.loadingManager.hide();
|
||||||
|
showToast('toast.agent.enrichFailed', { error: error.message }, 'error');
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
sendLoraToWorkflow(replaceMode) {
|
sendLoraToWorkflow(replaceMode) {
|
||||||
const card = this.currentCard;
|
const card = this.currentCard;
|
||||||
const usageTips = JSON.parse(card.dataset.usage_tips || '{}');
|
const usageTips = JSON.parse(card.dataset.usage_tips || '{}');
|
||||||
const loraSyntax = buildLoraSyntax(card.dataset.file_name, usageTips);
|
const folder = card.dataset.folder || '';
|
||||||
|
const loraName = folder ? `${folder}/${card.dataset.file_name}` : card.dataset.file_name;
|
||||||
|
const loraSyntax = buildLoraSyntax(loraName, usageTips);
|
||||||
|
|
||||||
sendLoraToWorkflow(loraSyntax, replaceMode, 'lora');
|
sendLoraToWorkflow(loraSyntax, replaceMode, 'lora');
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -187,6 +187,74 @@ export const ModelContextMenuMixin = {
|
|||||||
setTimeout(() => urlInput.focus(), 50);
|
setTimeout(() => urlInput.focus(), 50);
|
||||||
},
|
},
|
||||||
|
|
||||||
|
// HuggingFace linking methods
|
||||||
|
showLinkHfModal() {
|
||||||
|
const filePath = this.currentCard.dataset.filepath;
|
||||||
|
if (!filePath) return;
|
||||||
|
|
||||||
|
const confirmBtn = document.getElementById('confirmLinkHfBtn');
|
||||||
|
const urlInput = document.getElementById('hfModelUrl');
|
||||||
|
const errorDiv = document.getElementById('hfModelUrlError');
|
||||||
|
|
||||||
|
if (this._boundLinkHfHandler) {
|
||||||
|
confirmBtn.removeEventListener('click', this._boundLinkHfHandler);
|
||||||
|
}
|
||||||
|
|
||||||
|
this._boundLinkHfHandler = async () => {
|
||||||
|
const hfUrl = urlInput.value.trim();
|
||||||
|
if (!hfUrl) {
|
||||||
|
errorDiv.textContent = 'Please enter a HuggingFace repository URL.';
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
const hfPattern = /^https?:\/\/huggingface\.co\/([^/]+\/[^/]+)\/?$/;
|
||||||
|
if (!hfPattern.test(hfUrl)) {
|
||||||
|
errorDiv.textContent = 'Invalid URL format. Expected: https://huggingface.co/user/repo';
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
errorDiv.textContent = '';
|
||||||
|
modalManager.closeModal('linkHfModal');
|
||||||
|
|
||||||
|
try {
|
||||||
|
state.loadingManager.showSimpleLoading('Linking to HuggingFace...');
|
||||||
|
|
||||||
|
const response = await fetch('/api/lm/set-hf-url', {
|
||||||
|
method: 'POST',
|
||||||
|
headers: { 'Content-Type': 'application/json' },
|
||||||
|
body: JSON.stringify({ file_path: filePath, hf_url: hfUrl }),
|
||||||
|
});
|
||||||
|
|
||||||
|
if (!response.ok) {
|
||||||
|
const errData = await response.json().catch(() => ({}));
|
||||||
|
throw new Error(errData.error || `Request failed: ${response.statusText}`);
|
||||||
|
}
|
||||||
|
|
||||||
|
const data = await response.json();
|
||||||
|
if (data.success) {
|
||||||
|
showToast('toast.contextMenu.linkHfSuccess', {}, 'success');
|
||||||
|
await this.resetAndReload();
|
||||||
|
} else {
|
||||||
|
throw new Error(data.error || 'Failed to link model');
|
||||||
|
}
|
||||||
|
} catch (error) {
|
||||||
|
console.error('Error linking model to HuggingFace:', error);
|
||||||
|
showToast('toast.contextMenu.linkHfFailed', { message: error.message }, 'error');
|
||||||
|
} finally {
|
||||||
|
state.loadingManager.hide();
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
confirmBtn.addEventListener('click', this._boundLinkHfHandler);
|
||||||
|
|
||||||
|
urlInput.value = '';
|
||||||
|
errorDiv.textContent = '';
|
||||||
|
|
||||||
|
modalManager.showModal('linkHfModal');
|
||||||
|
|
||||||
|
setTimeout(() => urlInput.focus(), 50);
|
||||||
|
},
|
||||||
|
|
||||||
extractModelVersionId(url) {
|
extractModelVersionId(url) {
|
||||||
return extractCivitaiModelUrlParts(url);
|
return extractCivitaiModelUrlParts(url);
|
||||||
},
|
},
|
||||||
@@ -295,6 +363,9 @@ export const ModelContextMenuMixin = {
|
|||||||
case 'relink-civitai':
|
case 'relink-civitai':
|
||||||
this.showRelinkCivitaiModal();
|
this.showRelinkCivitaiModal();
|
||||||
return true;
|
return true;
|
||||||
|
case 'link-hf':
|
||||||
|
this.showLinkHfModal();
|
||||||
|
return true;
|
||||||
case 'set-nsfw':
|
case 'set-nsfw':
|
||||||
this.showNSFWLevelSelector(null, null, this.currentCard);
|
this.showNSFWLevelSelector(null, null, this.currentCard);
|
||||||
return true;
|
return true;
|
||||||
|
|||||||
@@ -260,8 +260,9 @@ export class RecipeContextMenu extends BaseContextMenu {
|
|||||||
strength: lora.strength || 1.0,
|
strength: lora.strength || 1.0,
|
||||||
|
|
||||||
// Model identifiers
|
// Model identifiers
|
||||||
|
modelId: lora.modelId || lora.model_id || civitaiInfo.modelId,
|
||||||
hash: modelFile?.hashes?.SHA256?.toLowerCase() || lora.hash,
|
hash: modelFile?.hashes?.SHA256?.toLowerCase() || lora.hash,
|
||||||
modelVersionId: civitaiInfo.id || lora.modelVersionId,
|
id: civitaiInfo.id || lora.modelVersionId,
|
||||||
|
|
||||||
// Metadata
|
// Metadata
|
||||||
thumbnailUrl: civitaiInfo.images?.[0]?.url || '',
|
thumbnailUrl: civitaiInfo.images?.[0]?.url || '',
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user