Compare commits

..

67 Commits

Author SHA1 Message Date
Will Miao bf6a614e0d chore(release): bump version to v1.1.7 2026-07-13 22:18:16 +08:00
Will Miao feab01cd9c fix(preview): hide license icons for models without CivitAI metadata 2026-07-13 19:49:10 +08:00
Will Miao 966024e534 fix(registry): force re-registration on WS refresh to prevent timeout, demote empty-registry log to debug
- workflow_registry.js: add force param to refreshRegistry(), bypass fingerprint
  dedup when responding to lora_registry_refresh WS message. Without this, the
  backend's wait_for_all() times out after 0.5s because the frontend skips the
  register-nodes POST when the workflow fingerprint hasn't changed (common after
  ComfyUI restart with an empty or unchanged workflow).
- misc_handlers.py: demote 'No nodes registered after refresh' from WARNING to
  DEBUG — empty workflows are a normal operational state, not a warning-worthy
  condition.
2026-07-13 19:10:48 +08:00
Will Miao 2018722cc8 fix(registry): handle compound subgraph node IDs, add proactive node push from graph hooks
- Handle compound node IDs (e.g. "252:0") from expanded group subgraphs
  to fix 400 Bad Request on workflows with group nodes
- Frontend proactively pushes node data via afterConfigureGraph and
  LiteGraph hooks (onNodeAdded/onNodeRemoved/graphChanged), eliminating
  WebSocket round-trip latency for most "Send to Workflow" operations
- Add content-fingerprint dedup to skip duplicate register-nodes POSTs
- Fast-path cache returns immediately when tabs are registered (including
  0-node registrations), avoiding unnecessary WS refresh cycles
- Distinguish "Empty Registry" from other errors in standalone UI toast
- Reduce WS refresh timeout 2s→0.5s, add cooldown and lock to prevent
  concurrent refresh storms
- All [LM:Registry] logs at DEBUG level
2026-07-13 18:02:26 +08:00
Will Miao 9d85c2a44a fix(ui): prevent tags widget from auto-resizing in Vue mode when tags change 2026-07-13 14:55:40 +08:00
Will Miao 03dd047e62 fix(download): return 200 instead of 500 when user cancels download 2026-07-13 11:47:48 +08:00
Will Miao 86b547c1e0 fix(locales): add missing downloadStopped key to toast.downloads section 2026-07-13 11:35:48 +08:00
Will Miao bab9752c8b fix(download): close modal before progress overlay and fix downloadId ReferenceError on cancel 2026-07-13 11:29:47 +08:00
Will Miao 774cc1be86 fix(download): use file ID for exact match, add debug logging for multi-file selection (#1023)
- Frontend: send file.id in file_params, use null instead of hardcoded defaults
- Backend: priority matching (ID exact → primary → lenient metadata)
- Lenient metadata: only compare fields present on both sides (fixes GGUF size mismatch)
- Add debug logs at key points: entry, file_params received, match result, anomaly signals
2026-07-13 11:15:03 +08:00
Will Miao 234b73c8a2 feat(ui): add cancel button to download progress modal 2026-07-13 09:40:53 +08:00
Will Miao abd06c48f4 fix(settings): reject checkpoints↔unet path overlap in extra folder paths with inline error UI
Changes:
- Backend: _validate_folder_paths() now checks checkpoints↔unet overlap
  within the same library using os.path.realpath() for symlink resolution
- Backend: set() calls _validate_folder_paths() for both folder_paths and
  extra_folder_paths before writing
- Backend: extracted _normalize_path_set() helper to eliminate duplicated
  normalization logic
- Frontend: inline error display with red border + error message below the
  conflicting input, no save triggered
- Frontend: path normalization (strip trailing slash, lowercase) in pre-check
  to reduce false negatives vs backend realpath
- Frontend: asymmetric error UX — message only on the user-edited side,
  red border on the pre-existing conflict side
- CSS: has-error styles with hardcoded rgba fallback for older browsers
- i18n: checkpointUnetOverlap + checkpointUnetOverlapInline keys added to
  all 10 locale files
2026-07-13 08:22:40 +08:00
Will Miao 6ca411e4e4 fix(ui): make loras widget fixed-size with user-controlled node resize
Remove dynamic height calculation that auto-resized the node when
LoRAs are added or removed. The widget now stays at the size the user
sets via the node resize handle, scrolling when content overflows.

- Drop updateWidgetHeight() and hardcoded entry-count height math
- Set --comfy-widget-min-height once (200px) instead of recalculating
- In Vue mode: add contain:layout+size to break the ResizeObserver
  feedback loop that forced node growth with content (CSS via
  .lm-loras-container.lm-vue-node scoped to vueNodesMode only)
- Remove unused "Node 2.0: Maximum visible LoRA entries" setting
2026-07-12 22:35:58 +08:00
Will Miao 6470021e77 feat(settings): persist LORA_MANAGER_PORTABLE to settings.json on first use (#1018) 2026-07-12 09:32:30 +08:00
Will Miao 71658ab37b feat(settings): add LORA_MANAGER_PORTABLE env var for per-instance settings isolation (#1018) 2026-07-12 07:44:31 +08:00
Will Miao 4f016a8024 feat(fetch): skip CivArchive API for HuggingFace-sourced models
- Bulk refresh filter now excludes models with hf_url
- Individual refresh for HF models only checks CivitAI API
- CivArchive client validates model IDs before querying
2026-07-11 20:29:54 +08:00
Will Miao f362ed585b fix(preview): gracefully handle deleted preview files - image fallback, cache cleanup, quieter logs
- Add onerror handler on <img> previews to fallback to no-preview.png
- Fire async cache cleanup when preview file returns 404
- Add ModelCache.clear_preview_by_path() for safe stale-url removal
- Downgrade /api/lm/previews 404 log from warning to debug
2026-07-10 21:25:07 +08:00
Will Miao 196172624f fix(ui): allow autocomplete textarea resize in app mode (#1020) 2026-07-09 11:59:09 +08:00
Will Miao 316702b7ab fix(hf): allow subdirectory paths in HF resolve URLs, strip repo-internal dirs on save (#1019) 2026-07-09 09:18:38 +08:00
Will Miao a7625b009f fix(ui): also exit bulk mode after enrich-hf-llm-bulk completes 2026-07-07 20:31:16 +08:00
Will Miao 5d4a33c90d fix(hf): stop using realpath for download path construction, match CivitAI approach 2026-07-07 20:24:47 +08:00
Will Miao 041a6b8525 Revert "fix(hf): pass computed folder to _save_hf_metadata instead of re-deriving from paths"
This reverts commit 54b44131b6.
2026-07-07 20:13:20 +08:00
Will Miao 2638109ad6 feat(hf): add Link to HuggingFace feature with unified Link Model submenu
- Merge Relink to Civitai and new Link to HuggingFace into a single
  'Link Model' submenu with sub-options for each source
- Add POST /api/lm/set-hf-url endpoint to associate a model with a
  HuggingFace repo URL, saving hf_url to .metadata.json
- Add link_hf_modal.html for URL input, following relink-civitai pattern
- Use update_single_model_cache instead of add_model_to_cache to
  prevent duplicate cache entries after linking
- Remove os.path.realpath usage for consistency with relink-civitai
- Raise errors instead of silently falling back to LoRA scanner when
  model root cannot be determined
- Scope .input-group CSS rules to modal IDs to fix style conflicts
  with download-modal.css
- Add i18n keys across all 10 locales with translations for
  zh-CN, zh-TW, ja, ko, de, es, fr, he, ru
2026-07-07 20:04:47 +08:00
Will Miao b019326747 feat(ui): auto-exit bulk mode after all bulk operations complete 2026-07-06 18:51:33 +08:00
Will Miao 54b44131b6 fix(hf): pass computed folder to _save_hf_metadata instead of re-deriving from paths 2026-07-06 17:34:43 +08:00
Will Miao a1d948025c fix(hf): strip empty trainedWords from metadata JSON to keep sidecar clean 2026-07-06 16:49:51 +08:00
Will Miao a90b2514ba feat(ui): group HF batch files by repo with collapse/expand, fix nested scroll & collapse animation
- Group HF batch download files by repo with collapsible group headers
- Fix nested scrollbar conflict (inner scrollbar undraggable) by making batch-preview-list flex-fill
- Fix collapse animation glitch (items disappearing before container shrinks) by keeping expanded during max-height transition
- Visual polish: hover lift, backdrop-filter glass, design token alignment
- Remove redundant database icon from group header
- Guard transitionend handlers against rapid-click races
2026-07-06 16:36:26 +08:00
pixelpaws cb4ad27813 Merge pull request #1013 from willmiao/agent
Hugging Face model metadata AI enrichment
2026-07-06 12:21:19 +08:00
Will Miao 637831248b fix(agent): route WS error events through onError instead of dead onComplete branch 2026-07-06 12:18:17 +08:00
Will Miao 00228deaaa fix(download): retry on Civitai 429 rate limit instead of removing images from metadata
When Civitai returns 429 (Too Many Requests) during example image
downloads, the previous behavior treated all failures identically and
permanently removed the corresponding images from model metadata —
making them impossible to retry.

This commit adds:
- 429 detection + Retry-After header parsing in download_to_memory
- Exponential backoff retry (up to 3 attempts) in
  download_model_images_with_tracking
- Separate tracking of rate-limited vs permanently failed URLs
- rate_limited_models progress tracking persisted to disk
- Rate-limited models are NOT added to failed_models/processed_models
  so they are automatically retried on subsequent download runs
- Force mode clears failed_models when rate-limited images exist
2026-07-06 11:58:19 +08:00
Will Miao 2373edf73c feat(ui): load provider model catalog asynchronously to avoid blocking page render 2026-07-06 10:02:09 +08:00
Will Miao e0e1b804a7 fix(llm): require api_base for custom provider without preset default 2026-07-06 10:02:04 +08:00
Will Miao fecbe8241f fix(agent): use status= instead of status_code in json_response calls 2026-07-06 10:02:00 +08:00
Will Miao 5983eaa1ce refactor(llm): use catalog-based max_tokens, remove JSON retry, reduce Ollama num_ctx
- Parse limit.output from model catalog alongside model IDs
  for per-model max output token limits
- Use catalog lookup in chat_completion_json() to set max_tokens;
  fall back to 4096 for unknown models (e.g. local Ollama)
- Remove the JSON retry (response_format → plain text fallback);
  keep _try_salvage_json as last-resort for truncated responses
- Reduce Ollama num_ctx from 32768 to 8192 (sufficient for
  metadata enrichment, saves VRAM)
- Fix stale test comment referencing removed retry
2026-07-06 09:13:42 +08:00
Will Miao 07fa454f72 chore(tests): stop tracking HF enrichment baseline snapshots
Remove tests/enrich_hf_validation/baselines/ from git tracking
(.gitignore entry + git rm --cached). These contain README snapshots
from community HF repos that may include NSFW/sensitive content.

Local files are preserved on disk for offline reference.
2026-07-06 01:08:25 +08:00
Will Miao 4b5aa45379 chore(tests): update bash code block tests to match preserved-bash behavior
Commit 9a0d866b changed _strip_fenced_code_blocks to preserve bash/shell
code blocks (they carry CLI setup and trigger-word metadata signal).
Update the two affected tests to expect bash content in the output
instead of asserting it is stripped.

- Rename test_bash_code_block_stripped → test_bash_code_block_preserved
- Update assertions: expect 'pip install' in result
2026-07-06 01:02:04 +08:00
Will Miao 9a0d866be4 fix(agent): preserve bash/shell code blocks in readme_processor during README cleaning 2026-07-06 00:40:35 +08:00
Will Miao 308d8f71b8 feat(ui): gray out enrich-hf-llm when no hf_url, add backend fast-fail, rename labels across locales, reposition menu item 2026-07-06 00:34:18 +08:00
Will Miao d0e8938039 fix(agent): call _format_base_models via self. to prevent NameError
The bare call  inside _build_prompt_context
would raise NameError because class methods don't close over class-level
scope. Use  instead to trigger attribute lookup.

Update enrich_hf_metadata prompt.md clue locations for better LLM accuracy.
Update baseline report to v2 (mean 69.0, 46 models, +2.2pp vs baseline 71.1%).
Consolidate README snapshots into baselines/readmes/.
2026-07-06 00:10:30 +08:00
Will Miao 13ed898b6b chore(tests): add base_model ground truth mapping for all 46 test entries 2026-07-05 20:47:30 +08:00
Will Miao e1dfd1c2a6 chore(tests): add two test entries and their HF README snapshots 2026-07-05 20:45:01 +08:00
Will Miao e3e944911b refactor(agent): extract shared scanner iteration into _find_model_entry
_Previous_ _find_scanner_for_model and identify_model_type contained ~25 lines
of identical scanner-iteration + path-matching logic.  Factor it into
_find_model_entry() so a new scanner type or edge-case fix can't drift apart.
2026-07-05 18:03:57 +08:00
Will Miao 51c0135250 refactor(agent): rename agent_cli to metadata_ops, strip temp debug logs
- Rename py/agent_cli/ -> py/metadata_ops/ (module was never agent-related)
- Rename tests/agent_cli/ -> tests/metadata_ops/
- Remove 9 low-value/debug INFO log points across agent_handlers.py,
  agent_service.py, llm_service.py, and metadata_ops/__init__.py
- Keep LLM raw response at DEBUG level for diagnostics
- Consolidate per-model progress + LLM result into single concise
  log line with basename instead of full path
- Update package/class/method docstrings to clarify this is a
  pipeline infrastructure, not a true agent loop
2026-07-05 18:00:58 +08:00
Will Miao 7b19bbb14e fix(agent): preserve preview URLs for collection repo models with flat heading structure
Three-part fix for enrich_hf_metadata failing to extract correct preview_url
from HuggingFace collection repos where models share flat heading levels:

1. _strip_standalone_images() now converts <img> tags to markdown image
   syntax ![alt](src) instead of stripping the URL entirely, so the LLM
   can still extract preview URLs.

2. _extract_section() uses a line-count-based forward window (stopping at
   <a id> anchors) for non-heading matches, instead of stopping at the
   very next heading. This prevents same-level sub-headings (# Download,
   # Trigger, # Sample prompt within a single model section) from
   truncating the window before sample images are included.

3. Post-processor preview fallback now filters gallery images to the
   model-specific README section before falling back to the repo-wide
   first image.
2026-07-05 17:05:47 +08:00
Will Miao 5494a70f40 chore(tests): commit validation dataset and baseline reports into repo
Move the HF model list from ~/Documents/ into tests/enrich_hf_validation/test_data/
and commit the pipeline validation baseline artifacts (report.json,
preprocessing_audit.json, README snapshots) into baselines/.

Update config.py and run_validation.py defaults to use repo-relative paths
via os.path.dirname(__file__) instead of ~/Documents/ hardcode.

Originates from changes in 8fb00998 (validation pipeline audit).
2026-07-05 17:03:45 +08:00
Will Miao 26c9ade1c9 feat(agent): optimize base model prompt — grouped display, comprehensive mapping rules, filename inference
- agent_service._format_base_models: output bullet list instead of
  JSON array for cleaner LLM parsing
- prompt.md mapping section: replace 14-row HF→CivitAI table with
  compact rule set covering 14 mapping paths including new entries
  for HiDream-ai, OnomaAIResearch/Illustrious, ideogram-ai/ideogram,
  Tongyi-MAI/Z-Image-Turbo, and Wan-AI/Wan2.*
- base_model extraction instruction: add guidance to infer from
  model filename, YAML tags, and README body text when YAML
  frontmatter has no explicit base_model:
2026-07-05 15:45:17 +08:00
Will Miao 87db23825f feat(constants): add 12 new CivitAI base models from API, sync JS/Python abbreviations and categories 2026-07-05 11:44:53 +08:00
Will Miao 8fb00998a7 feat(agent): fix extract_relevant_section false positives, add validation pipeline audit
- extract_relevant_section: raise token threshold >3, verify anchor
  sections contain basename, require 2+ heading token overlaps, skip
  TOC-style headings (markdown links), verify heading section size
- metadata_constructor: parse repo_id,model_name.safetensors format
  so model_path basename matches real filename
- config: replace hardcoded SUPPORTED_BASE_MODELS with dynamic
  init_supported_base_models() using production list_base_models()
- preprocessing_auditor: new Phase 1.5 audit module — fetches each
  README, runs extract_relevant_section + clean_readme_for_llm,
  records stats and flags, saves raw READMEs for cross-reference
- run_validation: integrate audit phase, add --audit-only mode,
  add LLM config consistency check, add ComfyUI root to sys.path
- report_generator: add Preprocessing Audit and Config Warnings
  sections to both markdown and JSON reports
2026-07-05 11:18:48 +08:00
Will Miao dd3aa97d0a refactor(agent): rename md_to_html to readme_processor, fix section extraction, widget parsing, and list_base_models
- Rename md_to_html.py → readme_processor.py (file no longer just HTML conversion)
- _extract_section: include YAML frontmatter, use heading-level-aware forward
  walk (sub-headings under # are included), increase walk limit past 30 lines
- _is_heading: exclude </hN> closing tags from boundary detection
- _heading_level: new helper for heading-level-aware section matching
- css: yield 0 for heading like closing tags, was unexpectedly caught by _is_heading
- extract_gallery_images: fix YAML block scalar (text: >-) prompt extraction;
  use endswith instead of == to detect the block marker
- _strip_widget_section: add to clean_readme_for_llm (widget text is handled
  by post-processor, not needed in LLM prompt)
- _strip_standalone_images: keep markdown image URLs intact for LLM preview
  extraction (was stripping to alt text only)
- list_base_models: switch from scanner-cache aggregation to
  CivitaiBaseModelService.get_base_models() - always returns full list
- Ollama: add num_ctx=32768 to payload options so thinking models have room
  to both reason and produce output
- Add tests/agent_cli/test_readme_processor.py: 59 tests covering extraction,
  cleaning, section matching, heading detection
- Update existing tests for behavioral changes
2026-07-05 06:39:54 +08:00
Will Miao 8bee8f4069 fix(recipe): fallback to locate custom example image on disk by model hash and image id (#1012) 2026-07-04 18:40:34 +08:00
Will Miao 817fe21b3e fix(ui): read cfg_scale and clip_skip with snake_case fallback, pass custom image id for recipe creation (#1012) 2026-07-04 18:40:24 +08:00
Will Miao 905c37290f chore: update runtime logs to use 'LLM enrichment' instead of 'Agent skill'
- agent_handlers.py: 'Agent skill' -> 'LLM enrichment' in all log messages
- skill_registry.py: 'agent skills' -> 'prompt-based skills' in discovery log
- llm_service.py: docstring 'agent skills' -> 'LLM-based enrichment features'
2026-07-04 16:53:41 +08:00
Will Miao f7632a47f9 feat(agent): enrich_hf_metadata with per-model progress and in-place card update
- PostProcessor returns updates dict from enrich_hf_metadata
- AgentService includes updated_data per model in WebSocket progress events
- Convert preview_url to HTTP URL via config.get_preview_static_url()
- LoraContextMenu: showEnhancedProgress + updateSingleItem per model
- BulkContextMenu: same pattern, remove window.location.reload()
- Guard empty updated_data and clean up callbacks on HTTP error
2026-07-04 16:50:56 +08:00
Will Miao 646f1ddfb1 refactor(agent): align 'Agent' naming to 'AI/LLM' to match current implementation
- locales/en.json: 'Enrich Metadata (Agent)' -> 'Enrich Metadata (AI)'
- Rename SKILL.md -> prompt.md with backward compat in skill_registry.py
- JS context menu action IDs: enrich-hf-agent -> enrich-hf-llm
- HTML template data-action attributes synced to match
- docstring cleanup: 'agent skill' -> 'skill pipeline' / 'feature'
2026-07-04 14:06:50 +08:00
Will Miao 170c8068c5 feat(agent): enrich_hf_metadata — filename-aware section matching, preview extraction for markdown/HTML/widget, JSON salvage, instance_prompt fallback, and validation suite
- extract_relevant_section(): trim README to model-filename-matching section
  for collection repos (download link, anchor ID, heading strategies)
- _strip_standalone_images(): preserve markdown image URLs so LLM can
  extract preview_url; strip only HTML <img> tags
- extract_simple_markdown_images(): extract civitai.images from ![]() body
- extract_html_img_tags(): extract from <img src="..."> (deadman44-style)
- extract_gallery_images(): fix widget parser for YAML - output: dash prefix
- _is_heading: exclude </hN> closing tags from boundary detection
- _extract_section: start at matching heading when match IS a heading line
- _try_salvage_json(): recover truncated JSON (close braces/brackets in
  LIFO order, close unterminated strings, strip trailing commas)
- PostProcessor: store _llm_confidence, add instance_prompt YAML fallback
- agent_service: pass model_basename to prompt, trim README via
  extract_relevant_section before clean_readme_for_llm
- Add tests/enrich_hf_validation/ suite: 100-model pipeline with progress
  checkpoint/resume, per-field scoring, markdown+JSON reporting
- Fix evaluation_engine: read _llm_confidence (not _llm_response)
2026-07-04 12:00:15 +08:00
Will Miao 3494037d20 fix(download): pass proxy to aria2 for actual file transfers (#1010) 2026-07-04 11:07:18 +08:00
Will Miao a1fd4e150b feat(agent): optimize enrich_hf_metadata with README cleaning, Ollama native API, and expanded fields
- Add clean_readme_for_llm() to strip noise from README before LLM injection
- Keep widget section text (valuable tag signal) and unmarked code blocks (trigger words)
- Preserve standalone image alt text instead of removing entirely
- Switch Ollama to native /api/chat with think:false to fix empty content on thinking models
- Extract Sample Gallery table images and deduplicate with widget images
- Only strip code blocks with explicit language tags (bash)
- Add notes and usage_tips fields to SKILL.md output format and post-processor
- Clean up dead code, fix regex edge cases, remove double type annotation
2026-07-04 08:01:50 +08:00
Will Miao b22f09bd1d fix(standalone): load extra folder paths from library settings in standalone mode 2026-07-03 19:21:56 +08:00
Will Miao 4ed9169646 feat(ui): redesign AI Provider settings with provider presets and model catalog
- Replace hardcoded provider list with PROVIDER_PRESETS (OpenAI, Ollama,
  DeepSeek, Groq, OpenRouter, OpenCode Go, Custom)
- Load model lists from models.dev/api.json catalog at startup
- Add Combobox vanilla JS component for model/base-URL selection
- Fetch local Ollama models via live API instead of catalog
- Hide API key values from frontend (boolean-only llm_api_key_set)
- Add i18n translations for all 9+ locales
- Update snapshot tests for new response fields
2026-07-03 16:08:51 +08:00
Will Miao f06c60bd47 fix(agent): handle plain YAML scalar text in extract_gallery_images
Widget entries with unquoted multi-line YAML scalars (e.g. "text: two samurais...\n  continuation") were not parsed, leaving gallery image prompts empty. Add a third branch for plain scalar format alongside the existing quoted and >- folded block handlers.
2026-07-03 07:34:24 +08:00
Will Miao ee8250c26c feat(agent): extract HF widget gallery images into civitai.images with recommended dimensions
- Add extract_gallery_images() to parse YAML widget entries from README
  frontmatter, convert relative image URLs to absolute HF URLs, and
  build civitai.images-compatible entries with prompt metadata
- LLM now extracts recommended_width/recommended_height from README
  (e.g. "Best Dimensions"), used as gallery image dimensions
- extract_gallery_images() accepts default_width/height parameters,
  falling back to 512x512 when LLM provides no recommendation
- Frontend ShowcaseView.js: defensive NaN guard for 0 width/height
- post_processor: consistently merge civitai updates across triggers,
  description, and gallery blocks with distinct variable names
- SKILL.md: add recommended_width/recommended_height to output schema
- 62 tests pass, including gallery extraction and dimension tests
2026-07-03 07:07:19 +08:00
Will Miao 88349bf944 feat(agent): render HF README as HTML in modelDescription, move converter to skill-local module
- Add inline convert_readme_to_html() in new skill-local md_to_html.py
  (zero external deps, handles h1-h4/bold/italic/code/lists/tables/links/hr)
- Strip YAML frontmatter, <Gallery />, badge images, HTML comments pre-conversion
- Fix indented whitespace after lists being misidentified as code blocks
- Fix HTML double-escaping in _inline_md (each pattern escapes independently)
- LLM short_description → civitai.description ("About this version" sidebar)
- raw README HTML → modelDescription (description tab, always available offline)
- Pass full readme_content from agent_service to post_processor
- 51 tests for converter + 4 updated/added post-processor tests
2026-07-02 23:34:52 +08:00
Will Miao a8adcaf023 feat(agent): improve enrich_hf_metadata skill with priority_tags, preview_url fix, civitai.trainedWords
- Add identify_model_type() helper to determine lora/checkpoint/embedding
- Pass priority_tags from user settings to LLM prompt for tag relevance
- SKILL.md: instruct LLM to exclude technical/generic HF tags, cross-reference
  against priority_tags; forbid ['None'] placeholder for trigger words
- post_processor: fix preview_url not updated after download (now writes local
  .webp path to metadata); write trigger words to civitai.trainedWords instead
  of top-level; sanitize ['None']/'null'/'n/a' placeholder values to []
- download_preview() now returns str | None (local path) instead of bool
- Update tests for new return type and nested civitai.trainedWords structure
2026-07-02 22:14:44 +08:00
Will Miao 63785f82b5 refactor(agent): consolidate skill definition into single SKILL.md with YAML frontmatter
Merge skill.yaml (metadata) and prompt.md (prompt template) into a
single SKILL.md file with YAML frontmatter, matching the agent-skill
convention used by opencode and Claude Code.

- Add frontmatter parser (_parse_skill_file) to SkillRegistry
- Remove skill.yaml, prompt.md, empty skills/__init__.py
- Remove obsolete load_handler method
- Update tests for new format and cleaned-up fields
2026-07-02 21:29:02 +08:00
Will Miao cf898da193 feat(agent): add LLM-powered metadata enrichment system with AgentCLI and PostProcessor
Introduce an agent skill framework for LLM-driven metadata enrichment:

- AgentCLI (py/agent_cli/): in-process wrappers around internal services
  using standard relative imports, eliminating the need for sys.path hacks
- LLMService: centralized BYOK (bring-your-own-key) LLM client supporting
  OpenAI, Ollama, and custom OpenAI-compatible endpoints
- PostProcessor: deterministic engine that applies LLM output via AgentCLI
  (replaces old handler.py + _BASE_MODEL_ALIASES approach)
- SkillRegistry: filesystem-based skill discovery (skill.yaml + prompt.md)
- AgentService: orchestrates skill execution with WebSocket progress
- Frontend AgentManager: WebSocket listeners, skill execution, config UI
- Context menu entries (single + bulk) for "Enrich Metadata (Agent)"
- Settings UI for AI Provider configuration (BYOK)
- Full i18n support across 9 locales

Bug fixes found during review:
- aiohttp.web.json_response: status_code= -> status=
- settings_modal cancelEditApiKey: wrong argument position
- AgentManager.isLlmConfigured: allow Ollama without API key
- PostProcessor._merge_tags: lowercase all tags to match TagUpdateService
2026-07-02 21:27:01 +08:00
Will Miao 3c83e78d9f feat(ui): auto-newline after pasting URL in download and batch-import textareas
Extract auto-newline-on-paste logic into shared setupAutoNewlineOnPaste() utility in uiHelpers.js.
Apply it to both the Download modal (modelUrl) and Batch Import modal (batchUrlInput)
textarea, so users can paste multiple URLs in succession without manually pressing Enter.
2026-07-02 10:53:33 +08:00
Will Miao d7291f73c9 fix(download): recognize civitai.red and civitai.green URLs in batch download (#1003) 2026-07-02 10:28:03 +08:00
Will Miao fe90f7f9b1 feat(ui): add searchable base model dropdown with filename inference in model modal
Replace native <select> with a searchable dropdown that:
- Filters options as the user types
- Shows filename-inferred suggestions at the top in a "Suggested" section
- Supports keyboard navigation (ArrowUp/Down/Enter/Escape)
- Allows typing custom values not in the list
- Removes dead .base-model-selector CSS

Adds 3 new i18n keys (baseModelSearchPlaceholder, baseModelSuggested,
baseModelNoMatch) with translations for all 9 locales.
2026-07-01 14:31:08 +08:00
122 changed files with 34694 additions and 22267 deletions
+4
View File
@@ -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/
+301 -285
View File
File diff suppressed because it is too large Load Diff
+208
View File
@@ -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` |
+2197 -2143
View File
File diff suppressed because it is too large Load Diff
+2195 -2141
View File
File diff suppressed because it is too large Load Diff
+2197 -2143
View File
File diff suppressed because it is too large Load Diff
+2197 -2143
View File
File diff suppressed because it is too large Load Diff
+2197 -2143
View File
File diff suppressed because it is too large Load Diff
+2197 -2143
View File
File diff suppressed because it is too large Load Diff
+2197 -2143
View File
File diff suppressed because it is too large Load Diff
+2197 -2143
View File
File diff suppressed because it is too large Load Diff
+2197 -2143
View File
File diff suppressed because it is too large Load Diff
+2197 -2143
View File
File diff suppressed because it is too large Load Diff
+21 -4
View File
@@ -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
@@ -1380,4 +1381,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
+11
View File
@@ -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)
+233
View File
@@ -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
+113
View File
@@ -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()
+6 -1
View File
@@ -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,
+165
View File
@@ -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,
)
+138 -39
View File
@@ -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:
+181 -33
View File
@@ -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
@@ -1562,6 +1585,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 +1643,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:
@@ -3056,6 +3129,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 +3177,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 +3223,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,
@@ -3317,6 +3456,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 +3476,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,6 +3492,8 @@ 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,
@@ -3384,6 +3527,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,
+24 -3
View File
@@ -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:
@@ -1303,9 +1313,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.
+31
View File
@@ -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:
+25
View File
@@ -2218,6 +2218,31 @@ class RecipeManagementHandler:
"Failed to download image for recipe: %s", exc "Failed to download image for recipe: %s", exc
) )
# Fallback: try to locate a custom image on disk using model_hash + image id
if image_bytes is None:
image_id = image_data.get("id") or ""
if image_id and model_hash:
from ...utils.example_images_paths import get_model_folder
model_folder = get_model_folder(model_hash)
if model_folder and os.path.exists(model_folder):
for fname in os.listdir(model_folder):
if f"custom_{image_id}" in fname:
ext = os.path.splitext(fname)[1].lower()
if ext not in (".jpg", ".jpeg", ".png", ".webp", ".gif"):
continue
fpath = os.path.join(model_folder, fname)
if os.path.isfile(fpath):
try:
with open(fpath, "rb") as f:
image_bytes = f.read()
extension = ext
except Exception as exc:
self._logger.warning(
"Failed to read custom image file %s: %s",
fpath, exc,
)
break
prompt = ( prompt = (
(parsed.get("gen_params") or {}).get("prompt") or "" (parsed.get("gen_params") or {}).get("prompt") or ""
) )
+15
View File
@@ -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"),
@@ -101,6 +103,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"
),
) )
+3
View File
@@ -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,
) )
+27
View File
@@ -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",
]
+489
View File
@@ -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)
+336
View File
@@ -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 `![alt](url)` 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 ""
+45
View File
@@ -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
+210
View File
@@ -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 `![alt](url)` 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
+7
View File
@@ -201,6 +201,13 @@ class Aria2Downloader:
"auto-file-renaming": "false", "auto-file-renaming": "false",
"file-allocation": "none", "file-allocation": "none",
} }
# Pass proxy to aria2 so the actual file transfer goes through the
# same proxy used by the aiohttp-based URL resolution step above.
downloader = await get_downloader()
if downloader.proxy_url:
options["all-proxy"] = downloader.proxy_url
if request_headers: if request_headers:
options["header"] = [ options["header"] = [
f"{key}: {value}" for key, value in request_headers.items() f"{key}: {value}" for key, value in request_headers.items()
+14
View File
@@ -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
+24
View File
@@ -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",
], ],
} }
+62 -14
View File
@@ -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:
@@ -1421,14 +1428,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 +1467,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,6 +1510,13 @@ 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"}
+37
View File
@@ -46,6 +46,30 @@ def is_ssl_cert_verify_error(exc: BaseException) -> bool:
return "CERTIFICATE_VERIFY_FAILED" in str(exc) return "CERTIFICATE_VERIFY_FAILED" in str(exc)
def _parse_retry_after(value: str) -> int:
"""Parse a Retry-After header value into seconds.
Supports both integer seconds and HTTP-date formats.
Returns a default of 60 seconds on invalid/missing input.
"""
if not value or not value.strip():
return 60
value = value.strip()
try:
return max(1, int(value))
except ValueError:
pass
try:
parsed = parsedate_to_datetime(value)
now = datetime.now().astimezone()
delta = (parsed - now).total_seconds()
return max(1, int(delta))
except (ValueError, OverflowError, OSError):
return 60
@dataclass(frozen=True) @dataclass(frozen=True)
class DownloadProgress: class DownloadProgress:
"""Snapshot of a download transfer at a moment in time.""" """Snapshot of a download transfer at a moment in time."""
@@ -911,6 +935,19 @@ class Downloader:
elif response.status == 404: elif response.status == 404:
error_msg = "File not found" error_msg = "File not found"
return False, error_msg, None return False, error_msg, None
elif response.status == 429:
raw_retry_after = response.headers.get("Retry-After")
retry_after = _parse_retry_after(raw_retry_after or "")
if raw_retry_after:
logger.warning(
"Rate limited (429) for %s, Retry-After: %ss", url, retry_after
)
else:
logger.warning(
"Rate limited (429) for %s, no Retry-After header; defaulting to %ss",
url, retry_after,
)
return False, f"Rate limited (429), retry after {retry_after}s", None
else: else:
error_msg = f"Download failed with status {response.status}" error_msg = f"Download failed with status {response.status}"
return False, error_msg, None return False, error_msg, None
+18
View File
@@ -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
+695
View File
@@ -0,0 +1,695 @@
"""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
result = await self.chat_completion(
messages=messages,
model=model,
temperature=temperature,
response_format={"type": "json_object"},
max_tokens=effective_max,
)
content = result.get("content", "") or ""
if not content:
raise LLMResponseError(
"LLM returned empty content in json_object mode. "
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
+15 -1
View File
@@ -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
+22 -1
View File
@@ -337,4 +337,25 @@ class ModelCache:
else: else:
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
+76 -1
View File
@@ -107,6 +107,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 +152,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 +630,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 +698,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 +924,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 +1588,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 +1847,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
+13 -1
View File
@@ -226,9 +226,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",
] ]
) )
+81 -24
View File
@@ -72,6 +72,7 @@ class _DownloadProgress(dict):
refreshed_models=set(), refreshed_models=set(),
failed_models=set(), failed_models=set(),
reprocessed_models=set(), reprocessed_models=set(),
rate_limited_models=set(),
) )
def snapshot(self) -> dict: def snapshot(self) -> dict:
@@ -82,6 +83,7 @@ class _DownloadProgress(dict):
snapshot["refreshed_models"] = list(self["refreshed_models"]) snapshot["refreshed_models"] = list(self["refreshed_models"])
snapshot["failed_models"] = list(self["failed_models"]) snapshot["failed_models"] = list(self["failed_models"])
snapshot["reprocessed_models"] = list(self.get("reprocessed_models", set())) snapshot["reprocessed_models"] = list(self.get("reprocessed_models", set()))
snapshot["rate_limited_models"] = list(self.get("rate_limited_models", set()))
return snapshot return snapshot
@@ -153,13 +155,15 @@ class DownloadManager:
# Step 3: Load progress file (I/O operation, done outside lock) # Step 3: Load progress file (I/O operation, done outside lock)
processed_models = set() processed_models = set()
failed_models = set() failed_models = set()
rate_limited_models = set()
try: try:
progress_file, processed_models, failed_models = await self._load_progress_file(output_dir) progress_file, processed_models, failed_models, rate_limited_models = await self._load_progress_file(output_dir)
logger.debug( logger.debug(
"Loaded previous progress, %s models already processed, %s models marked as failed", "Loaded previous progress, %s models already processed, %s models marked as failed, %s models rate-limited",
len(processed_models), len(processed_models),
len(failed_models), len(failed_models),
len(rate_limited_models),
) )
except Exception as e: except Exception as e:
logger.error(f"Failed to load progress file: {e}") logger.error(f"Failed to load progress file: {e}")
@@ -175,6 +179,7 @@ class DownloadManager:
self._progress.reset() self._progress.reset()
self._progress["processed_models"] = processed_models self._progress["processed_models"] = processed_models
self._progress["failed_models"] = failed_models self._progress["failed_models"] = failed_models
self._progress["rate_limited_models"] = rate_limited_models
self._stop_requested = False self._stop_requested = False
self._progress["status"] = "running" self._progress["status"] = "running"
self._progress["start_time"] = time.time() self._progress["start_time"] = time.time()
@@ -242,8 +247,8 @@ class DownloadManager:
"status": self._progress.snapshot(), "status": self._progress.snapshot(),
} }
async def _load_progress_file(self, output_dir: str) -> tuple[str, set, set]: async def _load_progress_file(self, output_dir: str) -> tuple[str, set, set, set]:
"""Load progress file from disk. Returns (progress_file_path, processed_models, failed_models). """Load progress file from disk. Returns (progress_file_path, processed_models, failed_models, rate_limited_models).
This is a separate async method to allow running in executor to avoid blocking event loop. This is a separate async method to allow running in executor to avoid blocking event loop.
""" """
@@ -252,8 +257,12 @@ class DownloadManager:
None, self._load_progress_file_sync, output_dir None, self._load_progress_file_sync, output_dir
) )
def _load_progress_file_sync(self, output_dir: str) -> tuple[str, set, set]: def _load_progress_file_sync(self, output_dir: str) -> tuple[str, set, set, set]:
"""Synchronous implementation of progress file loading.""" """Synchronous implementation of progress file loading.
Returns:
tuple: (progress_file_path, processed_models, failed_models, rate_limited_models)
"""
progress_file = os.path.join(output_dir, ".download_progress.json") progress_file = os.path.join(output_dir, ".download_progress.json")
progress_source = progress_file progress_source = progress_file
@@ -289,6 +298,7 @@ class DownloadManager:
processed_models = set() processed_models = set()
failed_models = set() failed_models = set()
rate_limited_models = set()
if os.path.exists(progress_source): if os.path.exists(progress_source):
try: try:
@@ -296,11 +306,11 @@ class DownloadManager:
saved_progress = json.load(f) saved_progress = json.load(f)
processed_models = set(saved_progress.get("processed_models", [])) processed_models = set(saved_progress.get("processed_models", []))
failed_models = set(saved_progress.get("failed_models", [])) failed_models = set(saved_progress.get("failed_models", []))
rate_limited_models = set(saved_progress.get("rate_limited_models", []))
except Exception: except Exception:
# Return empty sets on error
pass pass
return progress_file, processed_models, failed_models return progress_file, processed_models, failed_models, rate_limited_models
def _load_progress_sets_sync(self, progress_file: str) -> tuple[set, set]: def _load_progress_sets_sync(self, progress_file: str) -> tuple[set, set]:
"""Load only the processed and failed model sets from progress file. """Load only the processed and failed model sets from progress file.
@@ -732,11 +742,13 @@ class DownloadManager:
success, success,
is_stale, is_stale,
failed_images, failed_images,
rate_limited_images,
) = await ExampleImagesProcessor.download_model_images_with_tracking( ) = await ExampleImagesProcessor.download_model_images_with_tracking(
model_hash, model_name, images, model_dir, optimize, downloader model_hash, model_name, images, model_dir, optimize, downloader
) )
failed_urls: Set[str] = set(failed_images) failed_urls: Set[str] = set(failed_images)
rate_limited_urls: Set[str] = set(rate_limited_images)
# If metadata is stale, try to refresh it # If metadata is stale, try to refresh it
if is_stale and model_hash not in self._progress["refreshed_models"]: if is_stale and model_hash not in self._progress["refreshed_models"]:
@@ -760,6 +772,7 @@ class DownloadManager:
success, success,
_, _,
additional_failed, additional_failed,
additional_rate_limited,
) = await ExampleImagesProcessor.download_model_images_with_tracking( ) = await ExampleImagesProcessor.download_model_images_with_tracking(
model_hash, model_hash,
model_name, model_name,
@@ -770,29 +783,50 @@ class DownloadManager:
) )
failed_urls.update(additional_failed) failed_urls.update(additional_failed)
rate_limited_urls.update(additional_rate_limited)
self._progress["refreshed_models"].add(model_hash) self._progress["refreshed_models"].add(model_hash)
if failed_urls: # Separate permanent failures from rate-limited ones
permanent_failures = failed_urls - rate_limited_urls
if permanent_failures:
await self._remove_failed_images_from_metadata( await self._remove_failed_images_from_metadata(
model_hash, model_hash,
model_name, model_name,
model_dir, model_dir,
failed_urls, permanent_failures,
scanner, scanner,
) )
if failed_urls: if rate_limited_urls:
self._progress["rate_limited_models"].add(model_hash)
logger.warning(
"%d example images for %s are rate-limited (429), will retry next time",
len(rate_limited_urls),
model_name,
)
# Clear failed_models so non-force runs can retry
if force and model_hash in self._progress["failed_models"]:
self._progress["failed_models"].discard(model_hash)
logger.info(
f"Removed {model_name} from failed_models after force retry with rate-limited images"
)
if rate_limited_urls:
# Don't mark as failed or fully processed — rate-limited
# images will be retried next time.
pass
elif permanent_failures:
self._progress["failed_models"].add(model_hash) self._progress["failed_models"].add(model_hash)
self._progress["processed_models"].add(model_hash) self._progress["processed_models"].add(model_hash)
logger.info( logger.info(
"Removed %s failed example images for %s", "Removed %s failed example images for %s",
len(failed_urls), len(permanent_failures),
model_name, model_name,
) )
elif success: elif success:
self._progress["processed_models"].add(model_hash) self._progress["processed_models"].add(model_hash)
# Remove from failed_models if force mode enabled and model was previously failed
if force and model_hash in self._progress["failed_models"]: if force and model_hash in self._progress["failed_models"]:
self._progress["failed_models"].discard(model_hash) self._progress["failed_models"].discard(model_hash)
logger.info( logger.info(
@@ -850,6 +884,7 @@ class DownloadManager:
"processed_models": list(self._progress["processed_models"]), "processed_models": list(self._progress["processed_models"]),
"refreshed_models": list(self._progress["refreshed_models"]), "refreshed_models": list(self._progress["refreshed_models"]),
"failed_models": list(self._progress["failed_models"]), "failed_models": list(self._progress["failed_models"]),
"rate_limited_models": list(self._progress.get("rate_limited_models", set())),
"completed": self._progress["completed"], "completed": self._progress["completed"],
"total": self._progress["total"], "total": self._progress["total"],
"last_update": time.time(), "last_update": time.time(),
@@ -1155,11 +1190,13 @@ class DownloadManager:
success, success,
is_stale, is_stale,
failed_images, failed_images,
rate_limited_images,
) = await ExampleImagesProcessor.download_model_images_with_tracking( ) = await ExampleImagesProcessor.download_model_images_with_tracking(
model_hash, model_name, images, model_dir, optimize, downloader model_hash, model_name, images, model_dir, optimize, downloader
) )
failed_urls: Set[str] = set(failed_images) failed_urls: Set[str] = set(failed_images)
rate_limited_urls: Set[str] = set(rate_limited_images)
# If metadata is stale, try to refresh it # If metadata is stale, try to refresh it
if is_stale and model_hash not in self._progress["refreshed_models"]: if is_stale and model_hash not in self._progress["refreshed_models"]:
@@ -1183,6 +1220,7 @@ class DownloadManager:
success, success,
_, _,
additional_failed_images, additional_failed_images,
additional_rate_limited,
) = await ExampleImagesProcessor.download_model_images_with_tracking( ) = await ExampleImagesProcessor.download_model_images_with_tracking(
model_hash, model_hash,
model_name, model_name,
@@ -1192,21 +1230,35 @@ class DownloadManager:
downloader, downloader,
) )
# Combine failed images from both attempts
failed_urls.update(additional_failed_images) failed_urls.update(additional_failed_images)
rate_limited_urls.update(additional_rate_limited)
self._progress["refreshed_models"].add(model_hash) self._progress["refreshed_models"].add(model_hash)
# For forced downloads, remove failed images from metadata # Separate permanent failures from rate-limited ones
if failed_urls: permanent_failures = failed_urls - rate_limited_urls
# Only remove permanently failed images from metadata
if permanent_failures:
await self._remove_failed_images_from_metadata( await self._remove_failed_images_from_metadata(
model_hash, model_name, model_dir, failed_urls, scanner model_hash, model_name, model_dir, permanent_failures, scanner
) )
# Mark as processed if rate_limited_urls:
if ( self._progress["rate_limited_models"].add(model_hash)
success or failed_urls logger.warning(
): # Mark as processed if we successfully downloaded some images or removed failed ones "%d example images for %s are rate-limited (429), will retry next time",
len(rate_limited_urls),
model_name,
)
# Mark as processed only when no rate-limited images remain
if rate_limited_urls:
pass
elif permanent_failures:
self._progress["processed_models"].add(model_hash)
self._progress["failed_models"].add(model_hash)
elif success:
self._progress["processed_models"].add(model_hash) self._progress["processed_models"].add(model_hash)
return True # Return True to indicate a remote download happened return True # Return True to indicate a remote download happened
@@ -1229,15 +1281,20 @@ class DownloadManager:
model_dir: str, model_dir: str,
failed_images: Iterable[str], failed_images: Iterable[str],
scanner, scanner,
error_type: str = "not_found",
) -> None: ) -> None:
"""Mark failed images in model metadata so they won't be retried.""" """Mark failed images in model metadata so they won't be retried.
Args:
error_type: Reason string stored in the image's ``downloadError`` field
(default ``"not_found"``).
"""
failed_set: Set[str] = {url for url in failed_images if url} failed_set: Set[str] = {url for url in failed_images if url}
if not failed_set: if not failed_set:
return return
try: try:
# Get current model data
model_data = await MetadataUpdater.get_updated_model(model_hash, scanner) model_data = await MetadataUpdater.get_updated_model(model_hash, scanner)
if not model_data: if not model_data:
logger.warning( logger.warning(
@@ -1268,7 +1325,7 @@ class DownloadManager:
continue continue
image["downloadFailed"] = True image["downloadFailed"] = True
image.setdefault("downloadError", "not_found") image.setdefault("downloadError", error_type)
logger.debug( logger.debug(
"Marked example image %s for %s as failed due to missing remote asset", "Marked example image %s for %s as failed due to missing remote asset",
image_url, image_url,
+92 -39
View File
@@ -1,3 +1,4 @@
import asyncio
import logging import logging
import os import os
import re import re
@@ -194,16 +195,22 @@ class ExampleImagesProcessor:
return model_success, False # (success, is_metadata_stale) return model_success, False # (success, is_metadata_stale)
@staticmethod
def _extract_retry_after(error_message: str) -> int:
if not error_message:
return 60
match = re.search(r"retry after (\d+)s", str(error_message))
if match:
return max(1, int(match.group(1)))
return 60
@staticmethod @staticmethod
async def download_model_images_with_tracking(model_hash, model_name, model_images, model_dir, optimize, downloader): async def download_model_images_with_tracking(model_hash, model_name, model_images, model_dir, optimize, downloader):
"""Download images for a single model with tracking of failed image URLs
Returns:
tuple: (success, is_stale_metadata, failed_images) - whether download was successful, whether metadata is stale, list of failed image URLs
"""
model_success = True model_success = True
failed_images = [] failed_images = []
rate_limited_images = []
any_successful_download = False
for i, image in enumerate(model_images): for i, image in enumerate(model_images):
image_url = image.get('url') image_url = image.get('url')
if not image_url: if not image_url:
@@ -221,64 +228,110 @@ class ExampleImagesProcessor:
original_url = image_url original_url = image_url
if optimize and 'civitai.com' in image_url: if optimize and 'civitai.com' in image_url:
image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url) image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url)
# Download the file first to determine the actual file type async def _attempt_download() -> tuple:
try: logger.debug("Downloading media file %s for %s", i, model_name)
logger.debug(f"Downloading media file {i} for {model_name}") return await downloader.download_to_memory(
# Download using the unified downloader with headers
success, content, headers = await downloader.download_to_memory(
image_url, image_url,
use_auth=False, # Example images don't need auth use_auth=False,
return_headers=True return_headers=True,
) )
try:
success, content, headers = await _attempt_download()
if success: if success:
# Determine file extension from content or headers
media_ext = ExampleImagesProcessor._get_file_extension_from_content_or_headers( media_ext = ExampleImagesProcessor._get_file_extension_from_content_or_headers(
content, headers, original_url, image.get("type") content, headers, original_url, image.get("type")
) )
# Check if the detected file type is supported
is_image = media_ext in SUPPORTED_MEDIA_EXTENSIONS['images'] is_image = media_ext in SUPPORTED_MEDIA_EXTENSIONS['images']
is_video = media_ext in SUPPORTED_MEDIA_EXTENSIONS['videos'] is_video = media_ext in SUPPORTED_MEDIA_EXTENSIONS['videos']
if not (is_image or is_video): if not (is_image or is_video):
logger.debug(f"Skipping unsupported file type: {media_ext}") logger.debug("Skipping unsupported file type: %s", media_ext)
continue continue
# Use 0-based indexing with the detected extension
save_filename = f"image_{i}{media_ext}" save_filename = f"image_{i}{media_ext}"
save_path = os.path.join(model_dir, save_filename) save_path = os.path.join(model_dir, save_filename)
# Check if already downloaded
if os.path.exists(save_path): if os.path.exists(save_path):
logger.debug(f"File already exists: {save_path}") logger.debug("File already exists: %s", save_path)
continue continue
# Save the file
with open(save_path, 'wb') as f: with open(save_path, 'wb') as f:
f.write(content) f.write(content)
any_successful_download = True
elif ExampleImagesProcessor._is_not_found_error(content): elif ExampleImagesProcessor._is_not_found_error(content):
error_msg = f"Failed to download file: {image_url}, status code: 404 - Model metadata might be stale" error_msg = f"Failed to download file: {image_url}, status code: 404 - Model metadata might be stale"
logger.warning(error_msg) logger.warning(error_msg)
model_success = False # Mark the model as failed due to 404 error model_success = False
failed_images.append(image_url) # Track failed URL failed_images.append(image_url)
# Return early to trigger metadata refresh attempt return False, True, failed_images, rate_limited_images
return False, True, failed_images # (success, is_metadata_stale, failed_images)
elif "Rate limited (429)" in str(content):
max_attempts = 3
for attempt in range(1, max_attempts + 1):
wait = ExampleImagesProcessor._extract_retry_after(str(content)) * (2 ** (attempt - 1))
logger.warning(
"Rate limited (429) for %s, retry %d/%d after %ds",
image_url, attempt, max_attempts, wait,
)
await asyncio.sleep(wait)
success, content, headers = await _attempt_download()
if success:
media_ext = ExampleImagesProcessor._get_file_extension_from_content_or_headers(
content, headers, original_url, image.get("type")
)
is_image = media_ext in SUPPORTED_MEDIA_EXTENSIONS['images']
is_video = media_ext in SUPPORTED_MEDIA_EXTENSIONS['videos']
if not (is_image or is_video):
logger.debug("Skipping unsupported file type: %s", media_ext)
break
save_filename = f"image_{i}{media_ext}"
save_path = os.path.join(model_dir, save_filename)
if os.path.exists(save_path):
logger.debug("File already exists: %s", save_path)
break
with open(save_path, 'wb') as f:
f.write(content)
any_successful_download = True
break
elif "Rate limited (429)" in str(content):
continue
elif ExampleImagesProcessor._is_not_found_error(content):
logger.warning("Failed to download file: %s, status code: 404", image_url)
model_success = False
failed_images.append(image_url)
break
else:
logger.warning("Failed to download file: %s, error: %s", image_url, content)
model_success = False
failed_images.append(image_url)
break
else:
logger.warning(
"Giving up on %s after %d retries due to rate limiting",
image_url, max_attempts,
)
rate_limited_images.append(image_url)
model_success = False
else: else:
error_msg = f"Failed to download file: {image_url}, error: {content}" error_msg = f"Failed to download file: {image_url}, error: {content}"
logger.warning(error_msg) logger.warning(error_msg)
model_success = False # Mark the model as failed model_success = False
failed_images.append(image_url) # Track failed URL failed_images.append(image_url)
except Exception as e: except Exception as e:
error_msg = f"Error downloading file {image_url}: {str(e)}" error_msg = f"Error downloading file {image_url}: {str(e)}"
logger.error(error_msg) logger.error(error_msg)
model_success = False # Mark the model as failed model_success = False
failed_images.append(image_url) # Track failed URL failed_images.append(image_url)
return model_success, False, failed_images # (success, is_metadata_stale, failed_images) return any_successful_download or model_success, False, failed_images, rate_limited_images
@staticmethod @staticmethod
async def process_local_examples(model_file_path, model_file_name, model_name, model_dir, optimize): async def process_local_examples(model_file_path, model_file_name, model_name, model_dir, optimize):
+6
View File
@@ -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"""
+6 -1
View File
@@ -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
+1 -1
View File
@@ -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.1.7"
license = {file = "LICENSE"} license = {file = "LICENSE"}
dependencies = [ dependencies = [
"aiohttp", "aiohttp",
+150 -5
View File
@@ -444,16 +444,161 @@
flex: 1; flex: 1;
} }
.base-model-selector { /* ── Base Model Search Dropdown ─────────────────────────────────────────── */
width: 100%;
padding: 3px 5px; .base-model-search-wrapper {
position: relative;
flex: 1;
min-width: 0;
z-index: 100;
}
.base-model-search-input-wrapper {
display: flex;
align-items: center;
background: var(--bg-color); background: var(--bg-color);
border: 1px solid var(--lora-accent); border: 1px solid var(--lora-accent);
border-radius: var(--border-radius-xs); border-radius: var(--border-radius-xs);
padding: 0 6px;
gap: 4px;
}
.base-model-search-input-wrapper .search-icon {
color: var(--text-color);
opacity: 0.45;
font-size: 12px;
flex-shrink: 0;
pointer-events: none;
/* Reset global .search-icon rules from search-filter.css */
position: static;
right: auto;
top: auto;
transform: none;
}
.base-model-search-input {
flex: 1;
background: transparent;
border: none;
outline: none;
color: var(--text-color); color: var(--text-color);
font-size: 0.9em; font-size: 0.9em;
outline: none; padding: 3px 0;
margin-right: var(--space-1); width: 100%;
min-width: 0;
}
.base-model-search-input::placeholder {
color: var(--text-color);
opacity: 0.35;
}
.base-model-dropdown {
position: absolute;
top: 100%;
left: -1px;
right: -1px;
max-height: 270px;
overflow-y: auto;
background: var(--bg-color);
border: 1px solid var(--lora-border);
border-top: none;
border-radius: 0 0 var(--border-radius-xs) var(--border-radius-xs);
box-shadow: 0 8px 24px rgba(0, 0, 0, 0.22);
z-index: 101;
}
[data-theme="dark"] .base-model-dropdown {
box-shadow: 0 8px 28px rgba(0, 0, 0, 0.5);
}
/* Dropdown scrollbar styling */
.base-model-dropdown::-webkit-scrollbar {
width: 6px;
}
.base-model-dropdown::-webkit-scrollbar-thumb {
background: var(--lora-border);
border-radius: 3px;
}
.base-model-dropdown::-webkit-scrollbar-track {
background: transparent;
}
/* Section */
.base-model-dropdown-section {
border-bottom: 1px solid var(--lora-border);
}
.base-model-dropdown-section:last-child {
border-bottom: none;
}
/* Section header */
.base-model-dropdown-header {
padding: 5px 10px;
font-size: 0.72em;
font-weight: 600;
text-transform: uppercase;
letter-spacing: 0.08em;
color: var(--text-color);
opacity: 0.5;
background: var(--surface-subtle);
position: sticky;
top: 0;
z-index: 1;
}
.base-model-dropdown-header.suggested-header {
color: var(--lora-accent);
opacity: 1;
background: oklch(from var(--lora-accent) l c h / 0.08);
}
.base-model-dropdown-header.suggested-header i {
margin-right: 4px;
font-size: 0.85em;
}
/* Dropdown items */
.base-model-dropdown-item {
padding: 5px 12px;
cursor: pointer;
font-size: 0.9em;
color: var(--text-color);
transition: background 0.1s;
white-space: nowrap;
overflow: hidden;
text-overflow: ellipsis;
}
.base-model-dropdown-item:hover {
background: oklch(from var(--lora-accent) l c h / 0.1);
}
.base-model-dropdown-item.active {
background: oklch(from var(--lora-accent) l c h / 0.16);
}
.base-model-dropdown-item.selected {
font-weight: 600;
}
.base-model-dropdown-item.selected::after {
content: '✓';
float: right;
color: var(--lora-accent);
margin-left: 8px;
}
/* Empty state */
.base-model-dropdown-empty {
padding: 18px 12px;
text-align: center;
color: var(--text-color);
opacity: 0.4;
font-size: 0.88em;
} }
.size-wrapper { .size-wrapper {
+6
View 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);
} }
+119 -4
View File
@@ -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;
}
+12
View File
@@ -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();
+394
View File
@@ -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,6 +87,68 @@ 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 || '{}');
@@ -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;
+1 -1
View File
@@ -358,7 +358,7 @@ class RecipeCard {
<div class="delete-preview"> <div class="delete-preview">
${isVideo ? ${isVideo ?
`<video src="${previewUrl}" controls muted loop playsinline style="max-width: 100%;"></video>` : `<video src="${previewUrl}" controls muted loop playsinline style="max-width: 100%;"></video>` :
`<img src="${previewUrl}" alt="${this.recipe.title}">` `<img src="${previewUrl}" alt="${this.recipe.title}" onerror="this.onerror=null; this.src='/loras_static/images/no-preview.png'">`
} }
</div> </div>
<div class="delete-info"> <div class="delete-info">
+2 -2
View File
@@ -757,7 +757,7 @@ class RecipeModal {
`<video class="thumbnail-video" autoplay loop muted playsinline> `<video class="thumbnail-video" autoplay loop muted playsinline>
<source src="${lora.preview_url}" type="video/mp4"> <source src="${lora.preview_url}" type="video/mp4">
</video>` : </video>` :
`<img src="${lora.preview_url || '/loras_static/images/no-preview.png'}" alt="LoRA preview">`; `<img src="${lora.preview_url || '/loras_static/images/no-preview.png'}" alt="LoRA preview" onerror="this.onerror=null; this.src='/loras_static/images/no-preview.png'">`;
let loraItemClass = 'recipe-lora-item'; let loraItemClass = 'recipe-lora-item';
if (existsLocally) { if (existsLocally) {
@@ -1606,7 +1606,7 @@ class RecipeModal {
<video class="thumbnail-video" autoplay loop muted playsinline> <video class="thumbnail-video" autoplay loop muted playsinline>
<source src="${previewUrl}" type="video/mp4"> <source src="${previewUrl}" type="video/mp4">
</video> </video>
` : `<img src="${previewUrl}" alt="Checkpoint preview">`; ` : `<img src="${previewUrl}" alt="Checkpoint preview" onerror="this.onerror=null; this.src='/loras_static/images/no-preview.png'">`;
const badge = existsLocally ? ` const badge = existsLocally ? `
<div class="local-badge"> <div class="local-badge">
+1 -1
View File
@@ -643,7 +643,7 @@ export function createModelCard(model, modelType) {
<div class="card-preview ${shouldBlur ? 'blurred' : ''}"> <div class="card-preview ${shouldBlur ? 'blurred' : ''}">
${isVideo ? ${isVideo ?
`<video ${videoAttrs.join(' ')} style="pointer-events: none;"></video>` : `<video ${videoAttrs.join(' ')} style="pointer-events: none;"></video>` :
`<img src="${versionedPreviewUrl}" alt="${model.model_name}">` `<img src="${versionedPreviewUrl}" alt="${model.model_name}" onerror="this.onerror=null; this.src='/loras_static/images/no-preview.png'">`
} }
<div class="card-header"> <div class="card-header">
${shouldBlur ? ${shouldBlur ?
+279 -76
View File
@@ -6,6 +6,72 @@
import { BASE_MODEL_CATEGORIES, getMergedBaseModels } from '../../utils/constants.js'; import { BASE_MODEL_CATEGORIES, getMergedBaseModels } from '../../utils/constants.js';
import { showToast } from '../../utils/uiHelpers.js'; import { showToast } from '../../utils/uiHelpers.js';
import { getModelApiClient } from '../../api/modelApiFactory.js'; import { getModelApiClient } from '../../api/modelApiFactory.js';
import { translate } from '../../utils/i18nHelpers.js';
// ── Filename-based base model inference ──────────────────────────────────────
// Rules are ordered by specificity — first match wins for dedup.
// Each rule checks the filename (lowercased) for a regex pattern and suggests
// the associated base model values.
const BASE_MODEL_FILENAME_RULES = [
{ pattern: /flux\.?\s*2\s*klein/i, models: ['Flux.2 Klein 9B', 'Flux.2 Klein 9B-base', 'Flux.2 Klein 4B', 'Flux.2 Klein 4B-base'] },
{ pattern: /flux\.?\s*2/i, models: ['Flux.2 D', 'Flux.2 Klein 9B', 'Flux.2 Klein 4B'] },
{ pattern: /flux\.?\s*1\s*(dev|d)\b/i, models: ['Flux.1 D'] },
{ pattern: /flux\.?\s*1\s*(schnell|s)\b/i, models: ['Flux.1 S'] },
{ pattern: /flux/i, models: ['Flux.1 D', 'Flux.1 S', 'Flux.2 D'] },
{ pattern: /sdxl/i, models: ['SDXL 1.0', 'SDXL Lightning', 'SDXL Hyper'] },
{ pattern: /sd\s*1[._-\s]?5/i, models: ['SD 1.5'] },
{ pattern: /sd\s*1[._-\s]?4/i, models: ['SD 1.4'] },
{ pattern: /sd\s*1/i, models: ['SD 1.5', 'SD 1.4', 'SD 1.5 LCM', 'SD 1.5 Hyper'] },
{ pattern: /sd\s*3[._-\s]?5/i, models: ['SD 3.5', 'SD 3.5 Medium', 'SD 3.5 Large', 'SD 3.5 Large Turbo'] },
{ pattern: /sd\s*3/i, models: ['SD 3', 'SD 3.5'] },
{ pattern: /wan\s*\.?\s*video/i, models: ['Wan Video', 'Wan Video 1.3B t2v', 'Wan Video 14B t2v', 'Wan Video 14B i2v 480p', 'Wan Video 14B i2v 720p'] },
{ pattern: /hunyuan\s*\.?\s*video/i, models: ['Hunyuan Video'] },
{ pattern: /ltxv/i, models: ['LTXV', 'LTXV2', 'LTXV 2.3'] },
{ pattern: /cogvideo/i, models: ['CogVideoX'] },
{ pattern: /pony/i, models: ['Pony', 'Pony V7'] },
{ pattern: /illustrious/i, models: ['Illustrious'] },
{ pattern: /noobai/i, models: ['NoobAI'] },
{ pattern: /pixart/i, models: ['PixArt a', 'PixArt E'] },
{ pattern: /aura\s*\.?\s*flow/i, models: ['AuraFlow'] },
{ pattern: /kolors/i, models: ['Kolors'] },
{ pattern: /hunyuan\s*1/i, models: ['Hunyuan 1'] },
{ pattern: /lumina/i, models: ['Lumina'] },
{ pattern: /hidream/i, models: ['HiDream'] },
{ pattern: /qwen/i, models: ['Qwen'] },
{ pattern: /chroma/i, models: ['Chroma'] },
{ pattern: /anima/i, models: ['Anima'] },
{ pattern: /sd\s*2[._-\s]?[01]/i, models: ['SD 2.0', 'SD 2.1'] },
{ pattern: /mochi/i, models: ['Mochi'] },
{ pattern: /svd/i, models: ['SVD'] },
{ pattern: /zimage/i, models: ['ZImageTurbo', 'ZImageBase'] },
{ pattern: /nucleus/i, models: ['Nucleus'] },
{ pattern: /krea/i, models: ['Flux.1 Krea', 'Krea 2'] },
{ pattern: /ernie/i, models: ['Ernie', 'Ernie Turbo'] },
];
/**
* Infer likely base model(s) from a filename + model name string.
* Returns a deduplicated array in match-priority order.
* @param {string} filename
* @returns {string[]}
*/
function inferBaseModelsFromFilename(filename) {
if (!filename || typeof filename !== 'string') return [];
const seen = new Set();
const results = [];
for (const rule of BASE_MODEL_FILENAME_RULES) {
if (rule.pattern.test(filename)) {
for (const model of rule.models) {
if (!seen.has(model)) {
seen.add(model);
results.push(model);
}
}
}
}
return results;
}
/** /**
* Resolve the active file path for the currently open model modal. * Resolve the active file path for the currently open model modal.
@@ -226,7 +292,9 @@ export function setupModelNameEditing(filePath) {
} }
/** /**
* Set up base model editing functionality * Set up base model editing functionality with searchable dropdown
* Shows filename-inferred suggestions at the top, supports keyboard navigation,
* and allows typing custom values.
* @param {string} filePath - File path * @param {string} filePath - File path
*/ */
export function setupBaseModelEditing(filePath) { export function setupBaseModelEditing(filePath) {
@@ -257,116 +325,251 @@ export function setupBaseModelEditing(filePath) {
// Store the original value to check for changes later // Store the original value to check for changes later
const originalValue = baseModelContent.textContent.trim(); const originalValue = baseModelContent.textContent.trim();
// Create dropdown selector to replace the base model content // ── Build the full option list ────────────────────────────────────────
const currentValue = originalValue; const allModels = []; // { value, label, category }
const dropdown = document.createElement('select');
dropdown.className = 'base-model-selector';
// Flag to track if a change was made
let valueChanged = false;
// Add options from BASE_MODEL_CATEGORIES constants
const baseModelCategories = BASE_MODEL_CATEGORIES;
const categorizedModels = new Set(); const categorizedModels = new Set();
// Create option groups for better organization Object.entries(BASE_MODEL_CATEGORIES).forEach(([category, models]) => {
Object.entries(baseModelCategories).forEach(([category, models]) => {
const group = document.createElement('optgroup');
group.label = category;
models.forEach(model => { models.forEach(model => {
const option = document.createElement('option'); allModels.push({ value: model, label: model, category });
option.value = model;
option.textContent = model;
if (model === currentValue) option.selected = true;
categorizedModels.add(model); categorizedModels.add(model);
group.appendChild(option);
}); });
dropdown.appendChild(group);
}); });
// Check for dynamic base models from API that aren't in any category
const mergedModels = getMergedBaseModels(); const mergedModels = getMergedBaseModels();
const uncategorizedModels = mergedModels.filter(model => !categorizedModels.has(model)); const uncategorizedModels = mergedModels.filter(model => !categorizedModels.has(model));
if (uncategorizedModels.length > 0) { if (uncategorizedModels.length > 0) {
const group = document.createElement('optgroup');
group.label = 'Other (API)';
uncategorizedModels.forEach(model => { uncategorizedModels.forEach(model => {
const option = document.createElement('option'); allModels.push({ value: model, label: model, category: 'Other (API)' });
option.value = model;
option.textContent = model;
if (model === currentValue) option.selected = true;
group.appendChild(option);
}); });
dropdown.appendChild(group);
} }
// Replace content with dropdown // ── Filename-based inference ──────────────────────────────────────────
baseModelContent.style.display = 'none'; const fileName = (document.querySelector('.file-name-content')?.textContent || '') + ' ' +
baseModelDisplay.insertBefore(dropdown, editBtn); (document.querySelector('.model-name-content')?.textContent || '');
const inferredModels = inferBaseModelsFromFilename(fileName);
const inferredSet = new Set(inferredModels);
// Hide edit button during editing // ── Build search widget DOM ───────────────────────────────────────────
editBtn.style.display = 'none'; const wrapper = document.createElement('div');
wrapper.className = 'base-model-search-wrapper';
// Focus the dropdown // Search input row
dropdown.focus(); const inputWrapper = document.createElement('div');
inputWrapper.className = 'base-model-search-input-wrapper';
const searchIcon = document.createElement('i');
searchIcon.className = 'fas fa-search search-icon';
searchIcon.setAttribute('aria-hidden', 'true');
inputWrapper.appendChild(searchIcon);
const searchInput = document.createElement('input');
searchInput.type = 'text';
searchInput.className = 'base-model-search-input';
searchInput.placeholder = translate('modals.model.metadata.baseModelSearchPlaceholder', {}, 'Search base model…');
searchInput.autocomplete = 'off';
searchInput.spellcheck = false;
inputWrapper.appendChild(searchInput);
wrapper.appendChild(inputWrapper);
// Handle dropdown change // Dropdown list
dropdown.addEventListener('change', function() { const dropdown = document.createElement('div');
const selectedModel = this.value; dropdown.className = 'base-model-dropdown';
baseModelContent.textContent = selectedModel; wrapper.appendChild(dropdown);
// ── Render ────────────────────────────────────────────────────────────
function renderDropdown(filterText) {
const lowerFilter = (filterText || '').toLowerCase().trim();
dropdown.innerHTML = '';
let hasVisibleItems = false;
const fragment = document.createDocumentFragment();
// Mark that a change was made if the value differs from original // 1. Suggested section (filename-inferred, filtered by search)
if (selectedModel !== originalValue) { let suggestedToShow = inferredModels;
valueChanged = true; if (lowerFilter) {
} else { suggestedToShow = inferredModels.filter(m =>
valueChanged = false; m.toLowerCase().includes(lowerFilter)
);
}
if (suggestedToShow.length > 0) {
const section = document.createElement('div');
section.className = 'base-model-dropdown-section';
const header = document.createElement('div');
header.className = 'base-model-dropdown-header suggested-header';
header.innerHTML = '<i class="fas fa-star" aria-hidden="true"></i> ' +
translate('modals.model.metadata.baseModelSuggested', {}, 'Suggested');
section.appendChild(header);
suggestedToShow.forEach(model => {
const item = document.createElement('div');
item.className = 'base-model-dropdown-item';
if (model === originalValue) item.classList.add('selected');
item.dataset.value = model;
item.textContent = model;
section.appendChild(item);
hasVisibleItems = true;
});
fragment.appendChild(section);
}
// 2. Categorized options (deduplicated against suggestions)
const categoryMap = {};
allModels.forEach(m => {
if (inferredSet.has(m.value)) return; // already shown in Suggested
if (lowerFilter && !m.label.toLowerCase().includes(lowerFilter)) return;
if (!categoryMap[m.category]) categoryMap[m.category] = [];
categoryMap[m.category].push(m);
});
Object.entries(categoryMap).forEach(([category, items]) => {
if (items.length === 0) return;
const section = document.createElement('div');
section.className = 'base-model-dropdown-section';
const header = document.createElement('div');
header.className = 'base-model-dropdown-header';
header.textContent = category;
section.appendChild(header);
items.forEach(m => {
const item = document.createElement('div');
item.className = 'base-model-dropdown-item';
if (m.value === originalValue) item.classList.add('selected');
item.dataset.value = m.value;
item.textContent = m.label;
section.appendChild(item);
hasVisibleItems = true;
});
fragment.appendChild(section);
});
// 3. Empty state
if (!hasVisibleItems) {
const empty = document.createElement('div');
empty.className = 'base-model-dropdown-empty';
empty.textContent = translate('modals.model.metadata.baseModelNoMatch', {}, 'No matching base models');
fragment.appendChild(empty);
}
dropdown.appendChild(fragment);
// Scroll the selected item into view
const selected = dropdown.querySelector('.base-model-dropdown-item.selected');
if (selected) {
selected.scrollIntoView({ block: 'nearest' });
}
}
// Initial render — show everything
renderDropdown('');
// ── Events ────────────────────────────────────────────────────────────
let filterTimeout;
searchInput.addEventListener('input', () => {
clearTimeout(filterTimeout);
filterTimeout = setTimeout(() => renderDropdown(searchInput.value), 50);
});
// Click to select
dropdown.addEventListener('click', (e) => {
const item = e.target.closest('.base-model-dropdown-item');
if (!item) return;
baseModelContent.textContent = item.dataset.value;
cleanup();
const finalValue = baseModelContent.textContent.trim();
if (finalValue !== originalValue) {
saveBaseModel(
getActiveModalFilePath(baseModelContent.dataset.filePath),
originalValue
);
} }
}); });
// Function to save changes and exit edit mode // Replace content with search widget
const saveAndExit = function() { baseModelContent.style.display = 'none';
// Check if dropdown still exists and remove it editBtn.style.display = 'none';
if (dropdown && dropdown.parentNode === baseModelDisplay) { baseModelDisplay.insertBefore(wrapper, editBtn);
baseModelDisplay.removeChild(dropdown); searchInput.focus();
// ── Cleanup ───────────────────────────────────────────────────────────
function cleanup() {
if (wrapper.parentNode === baseModelDisplay) {
baseModelDisplay.removeChild(wrapper);
} }
// Show the content and edit button
baseModelContent.style.display = ''; baseModelContent.style.display = '';
editBtn.style.display = ''; editBtn.style.display = '';
// Remove editing class
baseModelDisplay.classList.remove('editing'); baseModelDisplay.classList.remove('editing');
// Only save if the value has actually changed
if (valueChanged || baseModelContent.textContent.trim() !== originalValue) {
const resolvedPath = getActiveModalFilePath(baseModelContent.dataset.filePath);
saveBaseModel(resolvedPath, originalValue);
}
// Remove this event listener
document.removeEventListener('click', outsideClickHandler); document.removeEventListener('click', outsideClickHandler);
}; }
// Handle outside clicks to save and exit // Outside click save typed/custom value if any
const outsideClickHandler = function(e) { const outsideClickHandler = function(e) {
// If click is outside the dropdown and base model display if (wrapper.contains(e.target)) return;
if (!baseModelDisplay.contains(e.target)) {
saveAndExit(); // If user typed a custom value (not just empty), apply it
const typedValue = searchInput.value.trim();
if (typedValue) {
baseModelContent.textContent = typedValue;
}
cleanup();
const finalValue = baseModelContent.textContent.trim();
if (finalValue !== originalValue) {
saveBaseModel(
getActiveModalFilePath(baseModelContent.dataset.filePath),
originalValue
);
} }
}; };
// Add delayed event listener for outside clicks // Defer listener to avoid the opening click itself
setTimeout(() => { setTimeout(() => {
document.addEventListener('click', outsideClickHandler); document.addEventListener('click', outsideClickHandler);
}, 0); }, 0);
// Also handle dropdown blur event // Keyboard navigation
dropdown.addEventListener('blur', function(e) { searchInput.addEventListener('keydown', function onKeydown(e) {
// Only save if the related target is not the edit button or inside the baseModelDisplay const items = Array.from(dropdown.querySelectorAll('.base-model-dropdown-item'));
if (!baseModelDisplay.contains(e.relatedTarget)) { const activeIdx = items.findIndex(el => el.classList.contains('active'));
saveAndExit();
if (e.key === 'ArrowDown') {
e.preventDefault();
items.forEach(el => el.classList.remove('active'));
const next = Math.min(activeIdx + 1, items.length - 1);
if (items[next]) {
items[next].classList.add('active');
items[next].scrollIntoView({ block: 'nearest' });
}
} else if (e.key === 'ArrowUp') {
e.preventDefault();
items.forEach(el => el.classList.remove('active'));
const prev = Math.max(activeIdx - 1, 0);
if (items[prev]) {
items[prev].classList.add('active');
items[prev].scrollIntoView({ block: 'nearest' });
}
} else if (e.key === 'Enter') {
e.preventDefault();
const activeItem = items.find(el => el.classList.contains('active'));
if (activeItem) {
activeItem.click();
} else if (searchInput.value.trim()) {
// Custom value typed
baseModelContent.textContent = searchInput.value.trim();
cleanup();
const finalValue = baseModelContent.textContent.trim();
if (finalValue !== originalValue) {
saveBaseModel(
getActiveModalFilePath(baseModelContent.dataset.filePath),
originalValue
);
}
}
} else if (e.key === 'Escape') {
e.preventDefault();
baseModelContent.textContent = originalValue;
cleanup();
} }
}); });
}); });
@@ -432,7 +432,7 @@ function renderMediaMarkup(version) {
return ` return `
<div class="version-media"> <div class="version-media">
<img src="${escapeHtml(version.previewUrl)}" alt="${escapeHtml(version.name || 'preview')}"> <img src="${escapeHtml(version.previewUrl)}" alt="${escapeHtml(version.name || 'preview')}" onerror="this.onerror=null; this.src='/loras_static/images/no-preview.png'">
</div> </div>
`; `;
} }
@@ -586,6 +586,7 @@ export function initMediaControlHandlers(container) {
const imageMetaRaw = this.dataset.imageMeta; const imageMetaRaw = this.dataset.imageMeta;
const imageUrl = this.dataset.imageUrl; const imageUrl = this.dataset.imageUrl;
const imageNsfw = this.dataset.imageNsfw; const imageNsfw = this.dataset.imageNsfw;
const imgId = this.dataset.imgId || '';
const localPath = this.dataset.localPath || ''; const localPath = this.dataset.localPath || '';
const showcaseSection = this.closest('.showcase-section'); const showcaseSection = this.closest('.showcase-section');
const modelHash = showcaseSection ? showcaseSection.dataset.modelHash : ''; const modelHash = showcaseSection ? showcaseSection.dataset.modelHash : '';
@@ -613,6 +614,7 @@ export function initMediaControlHandlers(container) {
meta: imageMeta, meta: imageMeta,
url: imageUrl, url: imageUrl,
nsfwLevel: imageNsfw ? parseInt(imageNsfw, 10) : undefined, nsfwLevel: imageNsfw ? parseInt(imageNsfw, 10) : undefined,
id: imgId || undefined,
}, },
model_hash: modelHash, model_hash: modelHash,
model_name: modelName || modelHash, model_name: modelName || modelHash,
@@ -174,7 +174,10 @@ function renderMediaItem(img, index, exampleFiles) {
const localUrl = localFile ? localFile.path : ''; const localUrl = localFile ? localFile.path : '';
// Calculate appropriate aspect ratio // Calculate appropriate aspect ratio
const aspectRatio = (img.height / img.width) * 100; // Defensive fallback: 0 width/height → 4:3 default (prevents NaN layout)
const safeW = img.width || 4;
const safeH = img.height || 3;
const aspectRatio = (safeH / safeW) * 100;
const containerWidth = 800; // modal content maximum width const containerWidth = 800; // modal content maximum width
const minHeightPercent = 40; const minHeightPercent = 40;
const maxHeightPercent = (window.innerHeight * 0.6 / containerWidth) * 100; const maxHeightPercent = (window.innerHeight * 0.6 / containerWidth) * 100;
@@ -210,8 +213,8 @@ function renderMediaItem(img, index, exampleFiles) {
const model = meta.Model || ''; const model = meta.Model || '';
const steps = meta.steps || ''; const steps = meta.steps || '';
const sampler = meta.sampler || ''; const sampler = meta.sampler || '';
const cfgScale = meta.cfgScale || ''; const cfgScale = meta.cfg_scale || meta.cfgScale || '';
const clipSkip = meta.clipSkip || ''; const clipSkip = meta.clip_skip || meta.clipSkip || '';
// Check if we have any meaningful generation parameters // Check if we have any meaningful generation parameters
const hasParams = seed || model || steps || sampler || cfgScale || clipSkip; const hasParams = seed || model || steps || sampler || cfgScale || clipSkip;
@@ -242,6 +245,7 @@ function renderMediaItem(img, index, exampleFiles) {
data-image-url="${img.url || ''}" data-image-url="${img.url || ''}"
data-image-nsfw="${img.nsfwLevel ?? ''}" data-image-nsfw="${img.nsfwLevel ?? ''}"
data-image-id="${cdnImageId}" data-image-id="${cdnImageId}"
data-img-id="${img.id || ''}"
data-local-path="${localFile ? localFile.path : ''}"> data-local-path="${localFile ? localFile.path : ''}">
<i class="fas fa-book-open"></i> <i class="fas fa-book-open"></i>
</button> </button>
+1
View File
@@ -15,6 +15,7 @@ import { initTheme, initBackToTop } from './utils/uiHelpers.js';
import { initializeInfiniteScroll } from './utils/infiniteScroll.js'; import { initializeInfiniteScroll } from './utils/infiniteScroll.js';
import { i18n } from './i18n/index.js'; import { i18n } from './i18n/index.js';
import { onboardingManager } from './managers/OnboardingManager.js'; import { onboardingManager } from './managers/OnboardingManager.js';
import './components/Combobox.js';
import { BulkContextMenu } from './components/ContextMenu/BulkContextMenu.js'; import { BulkContextMenu } from './components/ContextMenu/BulkContextMenu.js';
import { createPageContextMenu, createGlobalContextMenu } from './components/ContextMenu/index.js'; import { createPageContextMenu, createGlobalContextMenu } from './components/ContextMenu/index.js';
import { initializeEventManagement } from './utils/eventManagementInit.js'; import { initializeEventManagement } from './utils/eventManagementInit.js';
+209
View File
@@ -0,0 +1,209 @@
/**
* AgentManager WebSocket listener for agent skill progress events.
*
* Connects to the generic WebSocket endpoint and filters for
* `type: "agent_progress"` messages. Dispatches progress and completion
* events to registered callbacks.
*/
class AgentManager {
constructor() {
this.websocket = null;
this.progressCallbacks = [];
this.completeCallbacks = [];
this.errorCallbacks = [];
this.connected = false;
}
/**
* Connect to the WebSocket endpoint for agent progress events.
* Safe to call multiple times won't reconnect if already connected.
*/
connect() {
if (this.connected && this.websocket?.readyState === WebSocket.OPEN) {
return;
}
const wsProtocol = window.location.protocol === 'https:' ? 'wss://' : 'ws://';
try {
this.websocket = new WebSocket(
`${wsProtocol}${window.location.host}/ws/fetch-progress`
);
} catch (e) {
console.error('AgentManager: Failed to create WebSocket:', e);
return;
}
this.websocket.onopen = () => {
this.connected = true;
console.debug('AgentManager: WebSocket connected');
};
this.websocket.onmessage = (event) => {
try {
const data = JSON.parse(event.data);
if (data.type !== 'agent_progress') return;
this._dispatch(data);
} catch (e) {
// Not JSON or wrong format — ignore
}
};
this.websocket.onerror = (error) => {
console.error('AgentManager: WebSocket error:', error);
this.connected = false;
};
this.websocket.onclose = () => {
this.connected = false;
console.debug('AgentManager: WebSocket closed');
};
}
/**
* Dispatch a parsed agent event to the appropriate callbacks.
* @param {Object} data - The parsed WebSocket message
*/
_dispatch(data) {
const { status, skill } = data;
if (status === 'error') {
this.errorCallbacks.forEach((cb) => {
try {
cb(data);
} catch (e) {
console.error('AgentManager error callback failed:', e);
}
});
return;
}
if (status === 'completed') {
this.completeCallbacks.forEach((cb) => {
try {
cb(data);
} catch (e) {
console.error('AgentManager complete callback failed:', e);
}
});
return;
}
// started, processing — general progress
this.progressCallbacks.forEach((cb) => {
try {
cb(data);
} catch (e) {
console.error('AgentManager progress callback failed:', e);
}
});
}
/**
* Register a callback for progress events (started, processing).
* @param {Function} callback - Receives the event data
*/
onProgress(callback) {
this.progressCallbacks.push(callback);
}
/**
* Register a callback for completion events.
* @param {Function} callback - Receives the event data
*/
onComplete(callback) {
this.completeCallbacks.push(callback);
}
/**
* Register a callback for error events.
* @param {Function} callback - Receives the event data
*/
onError(callback) {
this.errorCallbacks.push(callback);
}
/**
* Clear all registered callbacks.
*/
clearCallbacks() {
this.progressCallbacks = [];
this.completeCallbacks = [];
this.errorCallbacks = [];
}
/**
* Execute an agent skill on the provided model paths.
*
* @param {string} skillName - The skill to execute
* @param {string[]} modelPaths - Model file paths to process
* @returns {Promise<Object>} The response JSON
*/
async executeSkill(skillName, modelPaths) {
const response = await fetch(
`/api/lm/agent/execute/${encodeURIComponent(skillName)}`,
{
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ model_paths: modelPaths }),
}
);
if (!response.ok) {
const errorData = await response.json().catch(() => ({}));
throw new Error(
errorData.error || `HTTP ${response.status}: ${response.statusText}`
);
}
return response.json();
}
/**
* Check if the LLM provider is configured.
*
* Returns true when both an API key and a model name are set.
*
* @returns {Promise<boolean>}
*/
_readProviderRequiresKey(providerId) {
const script = document.getElementById('llmProviderPresets');
if (!script) return true; // safe default
try {
const presets = JSON.parse(script.textContent);
const preset = presets[providerId];
return preset ? preset.requires_key !== false : true;
} catch {
return true;
}
}
async isLlmConfigured() {
try {
const response = await fetch('/api/lm/settings');
if (!response.ok) return false;
const data = await response.json();
const provider = data.settings?.llm_provider;
const hasModel = !!data.settings?.llm_model;
const hasKey = !!(data.settings?.llm_api_key_set || data.settings?.llm_api_key);
const needsKey = this._readProviderRequiresKey(provider);
return hasModel && (hasKey || !needsKey);
} catch {
return false;
}
}
/**
* Get the list of available agent skills.
*
* @returns {Promise<Array>}
*/
async listSkills() {
const response = await fetch('/api/lm/agent/skills');
if (!response.ok) return [];
const data = await response.json();
return data.skills || [];
}
}
// Export as singleton
export const agentManager = new AgentManager();
+4 -1
View File
@@ -1,5 +1,5 @@
import { modalManager } from './ModalManager.js'; import { modalManager } from './ModalManager.js';
import { showToast } from '../utils/uiHelpers.js'; import { showToast, setupAutoNewlineOnPaste } from '../utils/uiHelpers.js';
import { translate } from '../utils/i18nHelpers.js'; import { translate } from '../utils/i18nHelpers.js';
import { WS_ENDPOINTS } from '../api/apiConfig.js'; import { WS_ENDPOINTS } from '../api/apiConfig.js';
import { getStorageItem, setStorageItem } from '../utils/storageHelpers.js'; import { getStorageItem, setStorageItem } from '../utils/storageHelpers.js';
@@ -43,6 +43,9 @@ export class BatchImportManager {
setStorageItem('batch_import_skip_no_metadata', e.target.checked); setStorageItem('batch_import_skip_no_metadata', e.target.checked);
}); });
} }
// Auto-append newline after pasting a URL in the batch URL input
setupAutoNewlineOnPaste('batchUrlInput');
} }
/** /**
+18 -3
View File
@@ -633,7 +633,7 @@ export class BulkManager {
filePaths.forEach(path => { filePaths.forEach(path => {
state.virtualScroller.removeItemByFilePath(path); state.virtualScroller.removeItemByFilePath(path);
}); });
this.clearSelection(); if (state.bulkMode) this.toggleBulkMode();
if (window.modelDuplicatesManager) { if (window.modelDuplicatesManager) {
window.modelDuplicatesManager.updateDuplicatesBadgeAfterRefresh(); window.modelDuplicatesManager.updateDuplicatesBadgeAfterRefresh();
@@ -763,8 +763,9 @@ export class BulkManager {
`Re-import complete: ${completed} re-imported, ${failed} failed` `Re-import complete: ${completed} re-imported, ${failed} failed`
); );
const { resetAndReload: recipeResetAndReload } = await import('../api/recipeApi.js'); const { resetAndReload: recipeResetAndReload } = await import('../api/recipeApi.js');
recipeResetAndReload(false, { preserveScroll: false });
this.clearSelection(); this.clearSelection();
if (state.bulkMode) this.toggleBulkMode();
recipeResetAndReload(false, { preserveScroll: false });
} else { } else {
state.loadingManager.hide(); state.loadingManager.hide();
showToast('toast.recipes.reimportBulkFailed', {}, 'error'); showToast('toast.recipes.reimportBulkFailed', {}, 'error');
@@ -829,7 +830,7 @@ export class BulkManager {
); );
} }
this.clearSelection(); if (state.bulkMode) this.toggleBulkMode();
} else { } else {
throw new Error(result.error || 'Bulk repair failed'); throw new Error(result.error || 'Bulk repair failed');
} }
@@ -874,6 +875,8 @@ export class BulkManager {
if (this.isStripVisible) { if (this.isStripVisible) {
this.updateThumbnailStrip(); this.updateThumbnailStrip();
} }
if (state.bulkMode) this.toggleBulkMode();
} }
} catch (error) { } catch (error) {
@@ -927,6 +930,7 @@ export class BulkManager {
showToast('toast.models.bulkUpdatesNone', { type: typeLabel }, 'info'); showToast('toast.models.bulkUpdatesNone', { type: typeLabel }, 'info');
} }
if (state.bulkMode) this.toggleBulkMode();
await resetAndReload(false); await resetAndReload(false);
} catch (error) { } catch (error) {
console.error('Error checking updates for selected models:', error); console.error('Error checking updates for selected models:', error);
@@ -1273,6 +1277,8 @@ export class BulkManager {
showToast(toastKey, { count: failCount }, 'warning'); showToast(toastKey, { count: failCount }, 'warning');
} }
if (state.bulkMode) this.toggleBulkMode();
} catch (error) { } catch (error) {
console.error('Error during bulk tag operation:', error); console.error('Error during bulk tag operation:', error);
const toastKey = mode === 'replace' ? 'toast.models.bulkTagsReplaceFailed' : 'toast.models.bulkTagsAddFailed'; const toastKey = mode === 'replace' ? 'toast.models.bulkTagsReplaceFailed' : 'toast.models.bulkTagsAddFailed';
@@ -1398,6 +1404,8 @@ export class BulkManager {
} else { } else {
showToast('toast.models.bulkFavoriteFailed', {}, 'error'); showToast('toast.models.bulkFavoriteFailed', {}, 'error');
} }
if (state.bulkMode) this.toggleBulkMode();
} }
/** /**
@@ -1526,6 +1534,8 @@ export class BulkManager {
showToast('toast.models.bulkContentRatingFailed', {}, 'error'); showToast('toast.models.bulkContentRatingFailed', {}, 'error');
} }
if (state.bulkMode) this.toggleBulkMode();
return successCount > 0; return successCount > 0;
} }
@@ -1580,6 +1590,8 @@ export class BulkManager {
} else { } else {
showToast('toast.models.skipMetadataRefreshFailed', {}, 'error'); showToast('toast.models.skipMetadataRefreshFailed', {}, 'error');
} }
if (state.bulkMode) this.toggleBulkMode();
} }
/** /**
@@ -1674,6 +1686,8 @@ export class BulkManager {
showToast('toast.models.bulkBaseModelUpdateFailed', {}, 'error'); showToast('toast.models.bulkBaseModelUpdateFailed', {}, 'error');
} }
if (state.bulkMode) this.toggleBulkMode();
} catch (error) { } catch (error) {
console.error('Error during bulk base model operation:', error); console.error('Error during bulk base model operation:', error);
showToast('toast.models.bulkBaseModelUpdateFailed', {}, 'error'); showToast('toast.models.bulkBaseModelUpdateFailed', {}, 'error');
@@ -1711,6 +1725,7 @@ export class BulkManager {
// Call the auto-organize method with selected file paths // Call the auto-organize method with selected file paths
await apiClient.autoOrganizeModels(filePaths); await apiClient.autoOrganizeModels(filePaths);
if (state.bulkMode) this.toggleBulkMode();
resetAndReload(true); resetAndReload(true);
} catch (error) { } catch (error) {
console.error('Error during bulk auto-organize:', error); console.error('Error during bulk auto-organize:', error);
@@ -196,6 +196,17 @@ export class BulkMissingLoraDownloadManager {
let completedDownloads = 0; let completedDownloads = 0;
let failedDownloads = 0; let failedDownloads = 0;
let currentLoraProgress = 0; let currentLoraProgress = 0;
let cancelled = false;
loadingManager.showCancelButton(async () => {
if (cancelled) return;
cancelled = true;
try {
await this.loraApiClient.cancelDownload(batchDownloadId);
} catch (e) {
console.error('Cancel request failed:', e);
}
});
// Set up WebSocket message handler // Set up WebSocket message handler
ws.onmessage = (event) => { ws.onmessage = (event) => {
@@ -207,6 +218,11 @@ export class BulkMissingLoraDownloadManager {
return; return;
} }
if (data.status === 'cancelled') {
cancelled = true;
return;
}
// Process progress updates // Process progress updates
if (data.status === 'progress' && data.download_id && data.download_id.startsWith(batchDownloadId)) { if (data.status === 'progress' && data.download_id && data.download_id.startsWith(batchDownloadId)) {
currentLoraProgress = data.progress; currentLoraProgress = data.progress;
@@ -249,6 +265,8 @@ export class BulkMissingLoraDownloadManager {
// Download each LoRA sequentially // Download each LoRA sequentially
for (let i = 0; i < lorasToDownload.length; i++) { for (let i = 0; i < lorasToDownload.length; i++) {
if (cancelled) break;
const lora = lorasToDownload[i]; const lora = lorasToDownload[i];
currentLoraProgress = 0; currentLoraProgress = 0;
@@ -275,11 +293,13 @@ export class BulkMissingLoraDownloadManager {
modelId, modelId,
versionId, versionId,
loraRoot, loraRoot,
'', // Empty relative path, use default paths '',
useDefaultPaths, useDefaultPaths,
batchDownloadId batchDownloadId
); );
if (cancelled) break;
if (!response.success) { if (!response.success) {
console.error(`Failed to download LoRA ${lora.name || lora.file_name}: ${response.error}`); console.error(`Failed to download LoRA ${lora.name || lora.file_name}: ${response.error}`);
failedDownloads++; failedDownloads++;
@@ -288,8 +308,10 @@ export class BulkMissingLoraDownloadManager {
updateProgress(100, completedDownloads, ''); updateProgress(100, completedDownloads, '');
} }
} catch (error) { } catch (error) {
console.error(`Error downloading LoRA ${lora.name || lora.file_name}:`, error); if (!cancelled) {
failedDownloads++; console.error(`Error downloading LoRA ${lora.name || lora.file_name}:`, error);
failedDownloads++;
}
} }
} }
@@ -300,7 +322,10 @@ export class BulkMissingLoraDownloadManager {
loadingManager.hide(); loadingManager.hide();
// Show completion message // Show completion message
if (failedDownloads === 0) { if (cancelled) {
showToast('toast.downloads.downloadStopped', {}, 'info',
`Download cancelled. ${completedDownloads} item(s) completed.`);
} else if (failedDownloads === 0) {
showToast('toast.loras.allDownloadSuccessful', { count: completedDownloads }, 'success'); showToast('toast.loras.allDownloadSuccessful', { count: completedDownloads }, 'success');
} else { } else {
showToast('toast.loras.downloadPartialSuccess', { showToast('toast.loras.downloadPartialSuccess', {
+290 -97
View File
@@ -1,5 +1,5 @@
import { modalManager } from './ModalManager.js'; import { modalManager } from './ModalManager.js';
import { showToast } from '../utils/uiHelpers.js'; import { showToast, setupAutoNewlineOnPaste } from '../utils/uiHelpers.js';
import { state } from '../state/index.js'; import { state } from '../state/index.js';
import { LoadingManager } from './LoadingManager.js'; import { LoadingManager } from './LoadingManager.js';
import { getModelApiClient, resetAndReload } from '../api/modelApiFactory.js'; import { getModelApiClient, resetAndReload } from '../api/modelApiFactory.js';
@@ -31,6 +31,7 @@ export class DownloadManager {
// HF download state // HF download state
this.hfRepoId = null; this.hfRepoId = null;
this.hfSelectedFiles = []; this.hfSelectedFiles = [];
this.hfRepoCollapsed = {};
this.loadingManager = new LoadingManager(); this.loadingManager = new LoadingManager();
this.folderTreeManager = new FolderTreeManager(); this.folderTreeManager = new FolderTreeManager();
@@ -107,7 +108,8 @@ export class DownloadManager {
// Default path toggle handler // Default path toggle handler
document.getElementById('useDefaultPath').addEventListener('change', this.handleToggleDefaultPath); document.getElementById('useDefaultPath').addEventListener('change', this.handleToggleDefaultPath);
// Auto-append newline after pasting a URL so users can paste multiple URLs in succession
setupAutoNewlineOnPaste('modelUrl');
} }
updateModalLabels() { updateModalLabels() {
@@ -173,6 +175,7 @@ export class DownloadManager {
// Reset HF state // Reset HF state
this.hfRepoId = null; this.hfRepoId = null;
this.hfSelectedFiles = []; this.hfSelectedFiles = [];
this.hfRepoCollapsed = {};
} }
async retrieveVersionsForModel(modelId, source = null) { async retrieveVersionsForModel(modelId, source = null) {
@@ -463,8 +466,8 @@ export class DownloadManager {
const trimmed = url.trim(); const trimmed = url.trim();
if (!trimmed) return null; if (!trimmed) return null;
// CivitAI // CivitAI — matches civitai.com, civitai.red, civitai.green, etc.
if (/civitai\.com\/models\//i.test(trimmed) || /civitaiarchive|civarchive/i.test(trimmed)) { if (/civitai\.(?:com|red|green)\/models\//i.test(trimmed) || /civitaiarchive|civarchive/i.test(trimmed)) {
// Will be parsed by existing CivitAI logic // Will be parsed by existing CivitAI logic
return { type: 'civitai' }; return { type: 'civitai' };
} }
@@ -725,14 +728,23 @@ export class DownloadManager {
confirmFileSelection() { confirmFileSelection() {
const selectedRadio = document.querySelector('#fileSelectionList input[type="radio"]:checked'); const selectedRadio = document.querySelector('#fileSelectionList input[type="radio"]:checked');
if (!selectedRadio) return; if (!selectedRadio) {
console.warn('[download] confirmFileSelection: no radio button checked');
return;
}
const version = this.currentVersion; const version = this.currentVersion;
if (!version) return; if (!version) {
console.warn('[download] confirmFileSelection: no currentVersion set');
return;
}
const modelFiles = (version.files || []).filter(f => f.type === 'Model' || f.type === 'UNet' || f.type === 'Diffusion Model'); const modelFiles = (version.files || []).filter(f => f.type === 'Model' || f.type === 'UNet' || f.type === 'Diffusion Model');
this.selectedFile = modelFiles.find(f => f.id.toString() === selectedRadio.value); this.selectedFile = modelFiles.find(f => f.id.toString() === selectedRadio.value);
console.log('[download] confirmFileSelection: selected file id=%s, name="%s", type="%s", metadata=%o',
this.selectedFile?.id, this.selectedFile?.name, this.selectedFile?.type, this.selectedFile?.metadata);
document.getElementById('fileSelectionStep').style.display = 'none'; document.getElementById('fileSelectionStep').style.display = 'none';
document.getElementById('locationStep').style.display = 'block'; document.getElementById('locationStep').style.display = 'block';
this.proceedToLocationContent(); this.proceedToLocationContent();
@@ -869,16 +881,26 @@ export class DownloadManager {
const displayName = versionName || `#${versionId}`; const displayName = versionName || `#${versionId}`;
let ws = null; let ws = null;
let updateProgress = () => { }; let updateProgress = () => { };
let cancelled = false;
const downloadId = Date.now().toString();
try { try {
this.loadingManager.restoreProgressBar(); this.loadingManager.restoreProgressBar();
updateProgress = this.loadingManager.showDownloadProgress(1); updateProgress = this.loadingManager.showDownloadProgress(1);
updateProgress(0, 0, displayName); updateProgress(0, 0, displayName);
const downloadId = Date.now().toString();
const wsProtocol = window.location.protocol === 'https:' ? 'wss://' : 'ws://'; const wsProtocol = window.location.protocol === 'https:' ? 'wss://' : 'ws://';
ws = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${downloadId}`); ws = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${downloadId}`);
this.loadingManager.showCancelButton(async () => {
if (cancelled) return;
cancelled = true;
try {
await this.apiClient.cancelDownload(downloadId);
} catch (e) {
console.error('Cancel request failed:', e);
}
});
ws.onmessage = event => { ws.onmessage = event => {
const data = JSON.parse(event.data); const data = JSON.parse(event.data);
@@ -887,6 +909,12 @@ export class DownloadManager {
return; return;
} }
if (data.status === 'cancelled') {
cancelled = true;
this.loadingManager.setStatus(translate('modals.download.status.cancelled', {}, 'Download cancelled'));
return;
}
if (data.status === 'progress' && data.download_id === downloadId) { if (data.status === 'progress' && data.download_id === downloadId) {
const metrics = { const metrics = {
bytesDownloaded: data.bytes_downloaded, bytesDownloaded: data.bytes_downloaded,
@@ -925,6 +953,10 @@ export class DownloadManager {
fileParams fileParams
); );
if (cancelled) {
return false;
}
if (response?.skipped) { if (response?.skipped) {
this.loadingManager.setStatus(translate('modals.download.status.finalizing')); this.loadingManager.setStatus(translate('modals.download.status.finalizing'));
updateProgress(100, 0, displayName); updateProgress(100, 0, displayName);
@@ -965,8 +997,12 @@ export class DownloadManager {
return true; return true;
} catch (error) { } catch (error) {
console.error('Failed to download model version:', error); if (cancelled) {
showToast('toast.downloads.downloadError', { message: error?.message }, 'error'); console.log('Download cancelled by user:', downloadId);
} else {
console.error('Failed to download model version:', error);
showToast('toast.downloads.downloadError', { message: error?.message }, 'error');
}
return false; return false;
} finally { } finally {
try { try {
@@ -986,16 +1022,33 @@ export class DownloadManager {
const totalFiles = this.hfSelectedFiles.length; const totalFiles = this.hfSelectedFiles.length;
const updateProgress = this.loadingManager.showDownloadProgress(totalFiles); const updateProgress = this.loadingManager.showDownloadProgress(totalFiles);
let cancelled = false;
let currentDownloadId = null;
this.loadingManager.showCancelButton(async () => {
if (cancelled) return;
cancelled = true;
if (currentDownloadId) {
try {
await this.apiClient.cancelDownload(currentDownloadId);
} catch (e) {
console.error('Cancel request failed:', e);
}
}
});
try { try {
let completedDownloads = 0; let completedDownloads = 0;
for (let i = 0; i < totalFiles; i++) { for (let i = 0; i < totalFiles; i++) {
if (cancelled) break;
const filename = this.hfSelectedFiles[i]; const filename = this.hfSelectedFiles[i];
updateProgress(0, completedDownloads, filename); updateProgress(0, completedDownloads, filename);
this.loadingManager.setStatus(`Downloading ${filename}...`); this.loadingManager.setStatus(`Downloading ${filename}...`);
const downloadId = Date.now().toString() + '_' + i; currentDownloadId = Date.now().toString() + '_' + i;
const wsProtocol = window.location.protocol === 'https:' ? 'wss://' : 'ws://'; const wsProtocol = window.location.protocol === 'https:' ? 'wss://' : 'ws://';
const ws = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${downloadId}`); const ws = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${currentDownloadId}`);
try { try {
await new Promise((resolve, reject) => { await new Promise((resolve, reject) => {
@@ -1003,12 +1056,13 @@ export class DownloadManager {
ws.onerror = reject; ws.onerror = reject;
}); });
// Capture completed count at WS creation time so progress
// updates arriving after completedDownloads increments still
// show the correct "N / total" position.
const snapshotCompleted = completedDownloads; const snapshotCompleted = completedDownloads;
ws.onmessage = (event) => { ws.onmessage = (event) => {
const data = JSON.parse(event.data); const data = JSON.parse(event.data);
if (data.status === 'cancelled') {
cancelled = true;
return;
}
if (data.status === 'progress') { if (data.status === 'progress') {
const metrics = { const metrics = {
bytesDownloaded: data.bytes_downloaded, bytesDownloaded: data.bytes_downloaded,
@@ -1026,9 +1080,11 @@ export class DownloadManager {
modelRoot, modelRoot,
relativePath: targetFolder, relativePath: targetFolder,
useDefaultPaths, useDefaultPaths,
download_id: downloadId, download_id: currentDownloadId,
}); });
if (cancelled) break;
if (response?.success) { if (response?.success) {
completedDownloads++; completedDownloads++;
updateProgress(100, completedDownloads, filename); updateProgress(100, completedDownloads, filename);
@@ -1038,13 +1094,19 @@ export class DownloadManager {
} }
} }
showToast('toast.loras.downloadCompleted', {}, 'success'); if (cancelled) {
// Reload page data — model is already in scanner cache via backend showToast('toast.downloads.downloadStopped', {}, 'info',
`Download cancelled. ${completedDownloads} item(s) completed.`);
} else {
showToast('toast.loras.downloadCompleted', {}, 'success');
}
await resetAndReload(true); await resetAndReload(true);
return true; return true;
} catch (error) { } catch (error) {
console.error('Failed to download HF model:', error); if (!cancelled) {
showToast('toast.downloads.downloadError', { message: error?.message }, 'error'); console.error('Failed to download HF model:', error);
showToast('toast.downloads.downloadError', { message: error?.message }, 'error');
}
return false; return false;
} finally { } finally {
this.loadingManager.hide(); this.loadingManager.hide();
@@ -1077,7 +1139,7 @@ export class DownloadManager {
showBatchPreviewStep() { showBatchPreviewStep() {
document.querySelectorAll('.download-step').forEach(step => step.style.display = 'none'); document.querySelectorAll('.download-step').forEach(step => step.style.display = 'none');
document.getElementById('batchPreviewStep').style.display = 'block'; document.getElementById('batchPreviewStep').style.display = 'flex';
const validCount = this.batchModels.filter(m => { const validCount = this.batchModels.filter(m => {
if (m.error) return false; if (m.error) return false;
@@ -1091,56 +1153,36 @@ export class DownloadManager {
const list = document.getElementById('batchPreviewList'); const list = document.getElementById('batchPreviewList');
const hasHfItems = this.batchModels.some(m => m.source === 'huggingface' && !m.error); const hasHfItems = this.batchModels.some(m => m.source === 'huggingface' && !m.error);
let itemsHtml = this.batchModels.map((item, index) => { // Error items render flat, outside any group
if (item.error) { const errorItemsHtml = this.batchModels.map((item, index) => {
return ` if (!item.error) return null;
<div class="batch-preview-item batch-preview-error" data-index="${index}"> return `
<div class="batch-preview-icon"> <div class="batch-preview-item batch-preview-error" data-index="${index}">
<i class="fas fa-exclamation-triangle"></i> <div class="batch-preview-icon">
</div> <i class="fas fa-exclamation-triangle"></i>
<div class="batch-preview-info">
<div class="batch-preview-name">${item.url}</div>
<div class="batch-preview-meta batch-preview-error-text">${item.error}</div>
</div>
<button class="batch-preview-remove" data-index="${index}" title="${translate('common.actions.remove', {}, 'Remove')}">
<i class="fas fa-times"></i>
</button>
</div> </div>
`; <div class="batch-preview-info">
} <div class="batch-preview-name">${item.url}</div>
<div class="batch-preview-meta batch-preview-error-text">${item.error}</div>
</div>
<button class="batch-preview-remove" data-index="${index}" title="${translate('common.actions.remove', {}, 'Remove')}">
<i class="fas fa-times"></i>
</button>
</div>
`;
}).filter(Boolean).join('');
// CivitAI items render flat, outside any group (unchanged)
const civitaiItemsHtml = this.batchModels.map((item, index) => {
if (item.error) return null;
if (item.source === 'huggingface') return null;
const ver = item.selectedVersion; const ver = item.selectedVersion;
// HF batch item rendering with checkbox
if (item.source === 'huggingface') {
const hfSize = item.fileSizeBytes
? formatFileSize(item.fileSizeBytes)
: '?';
return `
<div class="batch-preview-item" data-index="${index}">
<input type="checkbox" class="batch-preview-checkbox"
data-index="${index}" ${item.checked !== false ? 'checked' : ''} />
<div class="batch-preview-info">
<div class="batch-preview-name">${item.displayName || item.filename || `HF #${index}`} <span class="hf-badge">HF</span></div>
<div class="batch-preview-meta">
<span>${hfSize}</span>
<span>${item.repo || ''}</span>
</div>
</div>
<button class="batch-preview-remove" data-index="${index}" title="${translate('common.actions.remove', {}, 'Remove')}">
<i class="fas fa-times"></i>
</button>
</div>
`;
}
const firstImage = ver?.images?.find(img => !img.url.endsWith('.mp4')); const firstImage = ver?.images?.find(img => !img.url.endsWith('.mp4'));
const thumbnailUrl = firstImage ? firstImage.url : '/loras_static/images/no-preview.png'; const thumbnailUrl = firstImage ? firstImage.url : '/loras_static/images/no-preview.png';
const fileSize = ver?.modelSizeKB const fileSize = ver?.modelSizeKB
? (ver.modelSizeKB / 1024).toFixed(1) ? (ver.modelSizeKB / 1024).toFixed(1)
: (ver?.files?.[0]?.sizeKB ? (ver.files[0].sizeKB / 1024).toFixed(1) : '?'); : (ver?.files?.[0]?.sizeKB ? (ver.files[0].sizeKB / 1024).toFixed(1) : '?');
const existsLocally = ver?.existsLocally; const existsLocally = ver?.existsLocally;
return ` return `
<div class="batch-preview-item ${existsLocally ? 'batch-preview-local' : ''}" data-index="${index}"> <div class="batch-preview-item ${existsLocally ? 'batch-preview-local' : ''}" data-index="${index}">
<div class="batch-preview-thumbnail"> <div class="batch-preview-thumbnail">
@@ -1161,8 +1203,59 @@ export class DownloadManager {
` : ''} ` : ''}
</div> </div>
`; `;
}).filter(Boolean).join('');
// Group HF items by repo (data model stays flat — only rendering groups)
const hfGroups = {};
this.batchModels.forEach((item, index) => {
if (item.error || item.source !== 'huggingface') return;
const repo = item.repo || 'unknown';
if (!hfGroups[repo]) hfGroups[repo] = [];
hfGroups[repo].push({ item, index });
});
const renderHfItem = ({ item, index }) => {
const hfSize = item.fileSizeBytes ? formatFileSize(item.fileSizeBytes) : '?';
return `
<div class="batch-preview-item" data-index="${index}">
<input type="checkbox" class="batch-preview-checkbox"
data-index="${index}" ${item.checked !== false ? 'checked' : ''} />
<div class="batch-preview-info">
<div class="batch-preview-name">${item.displayName || item.filename || `HF #${index}`} <span class="hf-badge">HF</span></div>
<div class="batch-preview-meta">
<span>${hfSize}</span>
<span>${item.repo || ''}</span>
</div>
</div>
<button class="batch-preview-remove" data-index="${index}" title="${translate('common.actions.remove', {}, 'Remove')}">
<i class="fas fa-times"></i>
</button>
</div>
`;
};
const hfGroupsHtml = Object.keys(hfGroups).map(repo => {
const items = hfGroups[repo];
const isCollapsed = this.hfRepoCollapsed[repo] === true;
const allChecked = items.every(({ item }) => item.checked !== false);
const fileCount = items.length;
return `
<div class="batch-preview-group" data-repo="${repo}">
<div class="batch-preview-group-header">
<i class="fas fa-chevron-right batch-preview-group-toggle ${isCollapsed ? '' : 'expanded'}"></i>
<span class="batch-preview-group-name">${repo}</span>
<span class="batch-preview-group-count">${fileCount} ${translate('modals.download.fileSelection.files', {}, 'files')}</span>
<input type="checkbox" class="batch-preview-group-select-all" data-repo="${repo}" ${allChecked ? 'checked' : ''} />
</div>
<div class="batch-preview-group-body ${isCollapsed ? '' : 'expanded'}">
${items.map(renderHfItem).join('')}
</div>
</div>
`;
}).join(''); }).join('');
let itemsHtml = errorItemsHtml + civitaiItemsHtml + hfGroupsHtml;
// Prepend select-all toolbar if there are HF items with checkboxes // Prepend select-all toolbar if there are HF items with checkboxes
if (hasHfItems) { if (hasHfItems) {
const allChecked = this.batchModels const allChecked = this.batchModels
@@ -1178,7 +1271,90 @@ export class DownloadManager {
list.innerHTML = itemsHtml; list.innerHTML = itemsHtml;
const updateCountAndSelectAll = () => {
const checkedCount = this.batchModels.filter(
m => !m.error && m.checked !== false
).length;
document.getElementById('downloadModalTitle').textContent =
translate('modals.download.titleWithType', { type: this.apiClient.apiConfig.config.displayName }) +
` (${checkedCount})`;
const nextBtn = document.getElementById('nextFromBatchBtn');
nextBtn.disabled = checkedCount === 0;
nextBtn.classList.toggle('disabled', checkedCount === 0);
// Global select-all
const selectAll = document.getElementById('batchSelectAll');
if (selectAll) {
const hfItems = this.batchModels.filter(m => m.source === 'huggingface' && !m.error);
selectAll.checked = hfItems.length > 0 && hfItems.every(m => m.checked !== false);
}
// Per-group select-all
list.querySelectorAll('.batch-preview-group-select-all').forEach(gsa => {
const repo = gsa.dataset.repo;
const repoItems = this.batchModels.filter(m => m.source === 'huggingface' && !m.error && m.repo === repo);
gsa.checked = repoItems.length > 0 && repoItems.every(m => m.checked !== false);
});
};
list.onclick = (e) => { list.onclick = (e) => {
// Per-group select-all checkbox
const groupSelectAll = e.target.closest('.batch-preview-group-select-all');
if (groupSelectAll) {
const repo = groupSelectAll.dataset.repo;
const checked = groupSelectAll.checked;
this.batchModels.forEach((m, idx) => {
if (m.source === 'huggingface' && !m.error && m.repo === repo) {
m.checked = checked;
const cb = list.querySelector(`.batch-preview-checkbox[data-index="${idx}"]`);
if (cb) cb.checked = checked;
}
});
updateCountAndSelectAll();
return;
}
const header = e.target.closest('.batch-preview-group-header');
if (header) {
const group = header.closest('.batch-preview-group');
const repo = group.dataset.repo;
const body = group.querySelector('.batch-preview-group-body');
const toggle = group.querySelector('.batch-preview-group-toggle');
const isCollapsed = this.hfRepoCollapsed[repo];
if (isCollapsed) {
this.hfRepoCollapsed[repo] = false;
body.style.transition = ''; // restore in case collapse was interrupted
body.classList.add('expanded');
toggle.classList.add('expanded');
// force reflow so expanded class is registered before setting height
void body.offsetHeight;
body.style.maxHeight = body.scrollHeight + 'px';
const onEnd = (e) => {
if (e.propertyName !== 'max-height') return;
if (this.hfRepoCollapsed[repo] !== false) return;
body.style.maxHeight = ''; // fall back to .expanded's 9999px
body.removeEventListener('transitionend', onEnd);
};
body.addEventListener('transitionend', onEnd);
} else {
this.hfRepoCollapsed[repo] = true;
body.style.maxHeight = body.scrollHeight + 'px';
requestAnimationFrame(() => {
// animate only max-height; keep expanded so opacity stays 1
body.style.transition = 'max-height 0.35s ease';
body.style.maxHeight = '0';
toggle.classList.remove('expanded');
const onEnd = (e) => {
if (e.propertyName !== 'max-height') return;
if (this.hfRepoCollapsed[repo] !== true) return; // state changed since
body.classList.remove('expanded');
body.style.transition = '';
body.removeEventListener('transitionend', onEnd);
};
body.addEventListener('transitionend', onEnd);
});
}
return;
}
const removeBtn = e.target.closest('.batch-preview-remove'); const removeBtn = e.target.closest('.batch-preview-remove');
if (removeBtn) { if (removeBtn) {
const idx = parseInt(removeBtn.dataset.index); const idx = parseInt(removeBtn.dataset.index);
@@ -1193,7 +1369,7 @@ export class DownloadManager {
} }
}; };
// Checkbox handler for HF batch items // Individual HF checkbox handler
const checkboxes = list.querySelectorAll('.batch-preview-checkbox'); const checkboxes = list.querySelectorAll('.batch-preview-checkbox');
checkboxes.forEach(cb => { checkboxes.forEach(cb => {
cb.addEventListener('change', (e) => { cb.addEventListener('change', (e) => {
@@ -1201,26 +1377,11 @@ export class DownloadManager {
if (this.batchModels[idx]) { if (this.batchModels[idx]) {
this.batchModels[idx].checked = e.target.checked; this.batchModels[idx].checked = e.target.checked;
} }
// Update valid count in title and Next button updateCountAndSelectAll();
const checkedCount = this.batchModels.filter(
m => !m.error && m.checked !== false
).length;
document.getElementById('downloadModalTitle').textContent =
translate('modals.download.titleWithType', { type: this.apiClient.apiConfig.config.displayName }) +
` (${checkedCount})`;
const nextBtn = document.getElementById('nextFromBatchBtn');
nextBtn.disabled = checkedCount === 0;
nextBtn.classList.toggle('disabled', checkedCount === 0);
// Update select-all checkbox state
const selectAll = document.getElementById('batchSelectAll');
if (selectAll) {
const hfItems = this.batchModels.filter(m => m.source === 'huggingface' && !m.error);
selectAll.checked = hfItems.length > 0 && hfItems.every(m => m.checked !== false);
}
}); });
}); });
// Select-all handler // Global select-all handler
const selectAll = document.getElementById('batchSelectAll'); const selectAll = document.getElementById('batchSelectAll');
if (selectAll) { if (selectAll) {
selectAll.addEventListener('change', (e) => { selectAll.addEventListener('change', (e) => {
@@ -1233,16 +1394,7 @@ export class DownloadManager {
this.batchModels[idx].checked = checked; this.batchModels[idx].checked = checked;
} }
}); });
// Update valid count in title and Next button updateCountAndSelectAll();
const checkedCount = this.batchModels.filter(
m => !m.error && m.checked !== false
).length;
document.getElementById('downloadModalTitle').textContent =
translate('modals.download.titleWithType', { type: this.apiClient.apiConfig.config.displayName }) +
` (${checkedCount})`;
const nextBtn = document.getElementById('nextFromBatchBtn');
nextBtn.disabled = checkedCount === 0;
nextBtn.classList.toggle('disabled', checkedCount === 0);
}); });
} }
@@ -1333,12 +1485,23 @@ export class DownloadManager {
} }
const fileParams = this.selectedFile ? { const fileParams = this.selectedFile ? {
id: this.selectedFile.id,
type: this.selectedFile.type || 'Model', type: this.selectedFile.type || 'Model',
format: this.selectedFile.metadata?.format || 'SafeTensor', format: this.selectedFile.metadata?.format || null,
size: this.selectedFile.metadata?.size || 'full', size: this.selectedFile.metadata?.size || null,
fp: this.selectedFile.metadata?.fp, fp: this.selectedFile.metadata?.fp || null,
} : null; } : null;
if (fileParams) {
console.log('[download] startDownload (single): fileParams built from selectedFile — id=%s, type=%s, format=%s, size=%s, fp=%s',
fileParams.id, fileParams.type, fileParams.format, fileParams.size, fileParams.fp);
} else {
console.log('[download] startDownload (single): this.selectedFile is null — no file selection, will download primary/default file. version=%s has %d files',
this.currentVersion?.id, (this.currentVersion?.files || []).length);
}
modalManager.closeModal('downloadModal');
return this.executeDownloadWithProgress({ return this.executeDownloadWithProgress({
modelId: this.modelId, modelId: this.modelId,
versionId: this.currentVersion.id, versionId: this.currentVersion.id,
@@ -1377,11 +1540,27 @@ export class DownloadManager {
let completedDownloads = 0; let completedDownloads = 0;
let failedDownloads = 0; let failedDownloads = 0;
let cancelled = false;
loadingManager.showCancelButton(async () => {
if (cancelled) return;
cancelled = true;
try {
await this.apiClient.cancelDownload(batchDownloadId);
} catch (e) {
console.error('Cancel request failed:', e);
}
});
ws.onmessage = (event) => { ws.onmessage = (event) => {
const data = JSON.parse(event.data); const data = JSON.parse(event.data);
if (data.type === 'download_id') return; if (data.type === 'download_id') return;
if (data.status === 'cancelled') {
cancelled = true;
return;
}
if (data.status === 'progress' && data.download_id?.startsWith(batchDownloadId)) { if (data.status === 'progress' && data.download_id?.startsWith(batchDownloadId)) {
const current = downloadItems[completedDownloads + failedDownloads]; const current = downloadItems[completedDownloads + failedDownloads];
const name = current?.selectedVersion?.name || current?.displayName || current?.filename || `#${completedDownloads + failedDownloads + 1}`; const name = current?.selectedVersion?.name || current?.displayName || current?.filename || `#${completedDownloads + failedDownloads + 1}`;
@@ -1400,6 +1579,8 @@ export class DownloadManager {
}); });
for (let i = 0; i < downloadItems.length; i++) { for (let i = 0; i < downloadItems.length; i++) {
if (cancelled) break;
const item = downloadItems[i]; const item = downloadItems[i];
const name = item.displayName || item.filename || (item.selectedVersion?.name || `Model #${item.modelId}`); const name = item.displayName || item.filename || (item.selectedVersion?.name || `Model #${item.modelId}`);
const isHf = item.source === 'huggingface'; const isHf = item.source === 'huggingface';
@@ -1410,7 +1591,6 @@ export class DownloadManager {
try { try {
let response; let response;
if (isHf) { if (isHf) {
// Per-file WebSocket for real-time progress
const downloadId = Date.now().toString() + '_hf_' + i; const downloadId = Date.now().toString() + '_hf_' + i;
const wsHf = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${downloadId}`); const wsHf = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${downloadId}`);
try { try {
@@ -1444,6 +1624,8 @@ export class DownloadManager {
wsHf.close(); wsHf.close();
} }
} else { } else {
console.log('[download] batch download: fileParams NOT passed for modelId=%s, versionId=%s — backend will use primary file',
item.modelId, item.selectedVersion?.id);
response = await this.apiClient.downloadModel( response = await this.apiClient.downloadModel(
item.modelId, item.modelId,
item.selectedVersion.id, item.selectedVersion.id,
@@ -1455,6 +1637,8 @@ export class DownloadManager {
); );
} }
if (cancelled) break;
if (!response.success) { if (!response.success) {
failedDownloads++; failedDownloads++;
} else { } else {
@@ -1462,15 +1646,20 @@ export class DownloadManager {
updateProgress(100, completedDownloads, ''); updateProgress(100, completedDownloads, '');
} }
} catch (err) { } catch (err) {
console.error(`Failed to download ${name}:`, err); if (!cancelled) {
failedDownloads++; console.error(`Failed to download ${name}:`, err);
failedDownloads++;
}
} }
} }
ws.close(); ws.close();
loadingManager.hide(); loadingManager.hide();
if (failedDownloads === 0) { if (cancelled) {
showToast('toast.downloads.downloadStopped', {}, 'info',
`Download cancelled. ${completedDownloads} item(s) completed.`);
} else if (failedDownloads === 0) {
showToast('toast.loras.allDownloadSuccessful', { count: completedDownloads }, 'success'); showToast('toast.loras.allDownloadSuccessful', { count: completedDownloads }, 'success');
} else { } else {
showToast('toast.loras.downloadPartialSuccess', { showToast('toast.loras.downloadPartialSuccess', {
@@ -1488,6 +1677,10 @@ export class DownloadManager {
modelRoot = '', modelRoot = '',
targetFolder = '' targetFolder = ''
} = {}) { } = {}) {
console.warn('[download] downloadVersionWithDefaults: NO fileParams will be sent — backend will always use primary file. '
+ 'modelType=%s, modelId=%s, versionId=%s, versionName="%s"',
modelType, modelId, versionId, versionName);
try { try {
this.apiClient = getModelApiClient(modelType); this.apiClient = getModelApiClient(modelType);
} catch (error) { } catch (error) {
+4
View File
@@ -281,6 +281,10 @@ export class LoadingManager {
// Initialize transfer stats with empty data // Initialize transfer stats with empty data
updateTransferStats(); updateTransferStats();
if (this.cancelButton) {
this.loadingContent.appendChild(this.cancelButton);
}
// Return update function // Return update function
return (currentProgress, currentIndex = 0, currentName = '', metrics = {}) => { return (currentProgress, currentIndex = 0, currentName = '', metrics = {}) => {
// Update current item progress // Update current item progress
+13
View File
@@ -264,6 +264,19 @@ export class ModalManager {
}); });
} }
// Add linkHfModal registration
const linkHfModal = document.getElementById('linkHfModal');
if (linkHfModal) {
this.registerModal('linkHfModal', {
element: linkHfModal,
onClose: () => {
this.getModal('linkHfModal').element.style.display = 'none';
document.body.classList.remove('modal-open');
},
closeOnOutsideClick: true
});
}
// Add exampleAccessModal registration // Add exampleAccessModal registration
const exampleAccessModal = document.getElementById('exampleAccessModal'); const exampleAccessModal = document.getElementById('exampleAccessModal');
if (exampleAccessModal) { if (exampleAccessModal) {
+2 -1
View File
@@ -330,8 +330,9 @@ class MoveManager {
.filter(r => r.success) .filter(r => r.success)
.map(r => ({ original_file_path: r.original_file_path, new_file_path: r.new_file_path })); .map(r => ({ original_file_path: r.original_file_path, new_file_path: r.new_file_path }));
// Deselect moving items // Deselect moving items and exit bulk mode
this.bulkFilePaths.forEach(path => bulkManager.deselectItem(path)); this.bulkFilePaths.forEach(path => bulkManager.deselectItem(path));
if (state.bulkMode) bulkManager.toggleBulkMode();
} else { } else {
// Single move mode // Single move mode
const result = await apiClient.moveSingleModel(this.currentFilePath, targetPath, this.useDefaultPath); const result = await apiClient.moveSingleModel(this.currentFilePath, targetPath, this.useDefaultPath);
+263 -17
View File
@@ -789,6 +789,27 @@ export class SettingsManager {
} }
} }
async _fetchProviderModelsAsync() {
try {
const resp = await fetch('/api/lm/llm/provider-models');
if (!resp.ok) return;
const data = await resp.json();
if (data.success && data.models) {
this._providerModels = data.models;
// Refresh model combobox if the settings modal is still open.
// Skip when provider is Ollama — it fetches its own live list
// from the local Ollama API and we must not overwrite it.
const llmProviderSelect = document.getElementById('llmProvider');
const provider = llmProviderSelect ? llmProviderSelect.value : 'openai';
if (this._llmModelCombobox && provider !== 'ollama') {
this._llmModelCombobox.updatePresets(this._providerModels[provider] || []);
}
}
} catch (_) {
// Silently ignore — models stay empty until next modal open
}
}
async loadSettingsToUI() { async loadSettingsToUI() {
// Set frontend settings from state // Set frontend settings from state
const blurMatureContentCheckbox = document.getElementById('blurMatureContent'); const blurMatureContentCheckbox = document.getElementById('blurMatureContent');
@@ -827,6 +848,115 @@ export class SettingsManager {
// Update API key status display (do NOT pre-fill the input) // Update API key status display (do NOT pre-fill the input)
this.updateApiKeyStatus(); this.updateApiKeyStatus();
this.updateLlmApiKeyStatus();
// ── AI Provider settings ──────────────────────────────────────
// Load provider presets from the JSON script tag embedded in the template
this._providerPresets = {};
this._providerModels = {};
const presetsScript = document.getElementById('llmProviderPresets');
if (presetsScript) {
try {
this._providerPresets = JSON.parse(presetsScript.textContent);
} catch (_) {
this._providerPresets = {};
}
}
const modelsScript = document.getElementById('llmProviderModels');
if (modelsScript) {
try {
this._providerModels = JSON.parse(modelsScript.textContent);
} catch (_) {
this._providerModels = {};
}
}
// If the embedded provider models is empty (server did not block on
// the remote catalog during page render), fetch asynchronously.
if (!this._providerModels || Object.keys(this._providerModels).length === 0) {
this._fetchProviderModelsAsync();
}
const llmProviderSelect = document.getElementById('llmProvider');
if (llmProviderSelect) {
llmProviderSelect.value = state.global.settings.llm_provider || 'openai';
}
// Destroy previous combobox instances before creating new ones,
// since loadSettingsToUI() runs on every modal open.
if (this._llmApiBaseCombobox) { this._llmApiBaseCombobox.destroy(); }
if (this._llmModelCombobox) { this._llmModelCombobox.destroy(); }
const llmApiBaseInput = document.getElementById('llmApiBase');
if (llmApiBaseInput) {
llmApiBaseInput.value = state.global.settings.llm_api_base || '';
const presetUrls = Object.values(this._providerPresets)
.map(p => p.api_base)
.filter(Boolean);
if (typeof Combobox !== 'undefined') {
this._llmApiBaseCombobox = new Combobox(llmApiBaseInput, {
presets: presetUrls,
placeholder: 'https://api.openai.com/v1',
});
}
}
// Helper to update model Combobox presets from catalog / Ollama API
const llmModelInput = document.getElementById('llmModel');
this._llmModelCombobox = null;
if (llmModelInput && typeof Combobox !== 'undefined') {
const currentProvider = llmProviderSelect ? llmProviderSelect.value : 'openai';
const fallbackModels = currentProvider === 'ollama' ? [] : (this._providerModels[currentProvider] || []);
this._llmModelCombobox = new Combobox(llmModelInput, {
presets: fallbackModels,
placeholder: translate('settings.aiProvider.modelPlaceholder', {}, 'Select a model...'),
onSelect: (value) => {
state.global.settings.llm_model = value;
this.saveSetting('llm_model', value)
.then(() => showToast('toast.settings.settingsUpdated', { setting: 'model' }, 'success'))
.catch(() => {});
},
});
}
const _loadModelPresets = async (provider) => {
if (!this._llmModelCombobox) return;
if (provider === 'ollama') {
try {
const apiBase = document.getElementById('llmApiBase')?.value?.trim() || 'http://localhost:11434/v1';
const resp = await fetch(`/api/lm/llm/models?provider=ollama&api_base=${encodeURIComponent(apiBase)}`);
if (resp.ok) {
const data = await resp.json();
if (data.success && Array.isArray(data.models)) {
this._llmModelCombobox.updatePresets(data.models);
return;
}
}
} catch (_) {}
this._llmModelCombobox.updatePresets([]);
} else {
this._llmModelCombobox.updatePresets(this._providerModels[provider] || []);
}
};
_loadModelPresets(llmProviderSelect ? llmProviderSelect.value : 'openai');
// Provider change → auto-fill API Base URL + update model presets
if (llmProviderSelect) {
llmProviderSelect.addEventListener('change', () => {
const provider = llmProviderSelect.value;
const preset = this._providerPresets[provider];
if (preset) {
if (llmApiBaseInput && preset.api_base) {
llmApiBaseInput.value = preset.api_base;
if (this._llmApiBaseCombobox) {
this._llmApiBaseCombobox.setValue(preset.api_base);
}
llmApiBaseInput.dispatchEvent(new Event('blur'));
}
}
_loadModelPresets(provider);
});
}
const civitaiHostSelect = document.getElementById('civitaiHost'); const civitaiHostSelect = document.getElementById('civitaiHost');
if (civitaiHostSelect) { if (civitaiHostSelect) {
@@ -1563,13 +1693,15 @@ export class SettingsManager {
<input type="text" class="extra-folder-path-input" <input type="text" class="extra-folder-path-input"
placeholder="${translate('settings.extraFolderPaths.pathPlaceholder', {}, '/path/to/models')}" value="${path}" placeholder="${translate('settings.extraFolderPaths.pathPlaceholder', {}, '/path/to/models')}" value="${path}"
onblur="settingsManager.updateExtraFolderPaths('${modelType}')" onblur="settingsManager.updateExtraFolderPaths('${modelType}')"
onfocus="settingsManager.clearExtraFolderPathError(this)"
onkeydown="if(event.key === 'Enter') { this.blur(); }" /> onkeydown="if(event.key === 'Enter') { this.blur(); }" />
<button type="button" class="remove-path-btn" <button type="button" class="remove-path-btn"
onclick="this.parentElement.parentElement.remove(); settingsManager.updateExtraFolderPaths('${modelType}')" onclick="settingsManager.removeExtraFolderPathRow(this, '${modelType}')"
title="${translate('common.actions.delete', {}, 'Delete')}"> title="${translate('common.actions.delete', {}, 'Delete')}">
<i class="fas fa-times"></i> <i class="fas fa-times"></i>
</button> </button>
</div> </div>
<div class="extra-folder-path-error"></div>
`; `;
container.appendChild(row); container.appendChild(row);
@@ -1583,7 +1715,63 @@ export class SettingsManager {
} }
} }
clearExtraFolderPathError(input) {
input.classList.remove('has-error');
const row = input.closest('.extra-folder-path-row');
if (row) {
const errEl = row.querySelector('.extra-folder-path-error');
if (errEl) {
errEl.classList.remove('visible');
errEl.textContent = '';
}
}
}
_clearAllExtraFolderPathErrors() {
document.querySelectorAll('.extra-folder-path-input.has-error').forEach((input) => {
input.classList.remove('has-error');
});
document.querySelectorAll('.extra-folder-path-error.visible').forEach((el) => {
el.classList.remove('visible');
el.textContent = '';
});
}
_markExtraFolderPathsError(modelType, overlappingPaths, showMessage = false) {
const container = document.getElementById(`extraFolderPaths-${modelType}`);
if (!container) return;
const inputs = container.querySelectorAll('.extra-folder-path-input');
inputs.forEach((input) => {
const val = input.value.trim();
if (val && overlappingPaths.includes(val)) {
input.classList.add('has-error');
if (showMessage) {
const row = input.closest('.extra-folder-path-row');
if (row) {
const errEl = row.querySelector('.extra-folder-path-error');
if (errEl) {
errEl.textContent = translate('settings.extraFolderPaths.validation.checkpointUnetOverlapInline', {}, 'This path is also used for a different model type. Use separate folders for checkpoints and diffusion models.');
errEl.classList.add('visible');
}
}
}
}
});
}
removeExtraFolderPathRow(btn, modelType) {
const row = btn.closest('.extra-folder-path-row');
if (row) {
row.remove();
this.updateExtraFolderPaths(modelType);
}
}
async updateExtraFolderPaths(changedModelType) { async updateExtraFolderPaths(changedModelType) {
// Clear previous errors
this._clearAllExtraFolderPathErrors();
const extraFolderPaths = {}; const extraFolderPaths = {};
// Collect paths for all model types // Collect paths for all model types
@@ -1604,6 +1792,32 @@ export class SettingsManager {
extraFolderPaths[modelType] = paths; extraFolderPaths[modelType] = paths;
}); });
// Client-side pre-check: checkpoints and unet must not share the same path.
// Normalise paths to reduce false negatives vs the backend's realpath + normcase.
const normalise = (p) => p.replace(/[/\\]+$/, '').toLowerCase();
const ckptSet = new Set((extraFolderPaths.checkpoints || []).map(normalise));
const unetSet = new Set((extraFolderPaths.unet || []).map(normalise));
const ckptOverlap = (extraFolderPaths.checkpoints || []).filter(p => p && unetSet.has(normalise(p)));
const unetOverlap = (extraFolderPaths.unet || []).filter(p => p && ckptSet.has(normalise(p)));
const hasOverlap = ckptOverlap.length > 0 || unetOverlap.length > 0;
if (hasOverlap) {
// Error message only on the side the user just edited.
// The other side gets red border only (passive conflict indicator).
if (changedModelType === 'checkpoints') {
this._markExtraFolderPathsError('checkpoints', ckptOverlap, true);
this._markExtraFolderPathsError('unet', unetOverlap, false);
} else if (changedModelType === 'unet') {
this._markExtraFolderPathsError('unet', unetOverlap, true);
this._markExtraFolderPathsError('checkpoints', ckptOverlap, false);
} else {
// Pre-existing conflict from direct config edit — mark both without messages
this._markExtraFolderPathsError('checkpoints', ckptOverlap, false);
this._markExtraFolderPathsError('unet', unetOverlap, false);
}
return;
}
// Check if paths have actually changed // Check if paths have actually changed
const currentPaths = state.global.settings.extra_folder_paths || {}; const currentPaths = state.global.settings.extra_folder_paths || {};
const pathsChanged = JSON.stringify(currentPaths) !== JSON.stringify(extraFolderPaths); const pathsChanged = JSON.stringify(currentPaths) !== JSON.stringify(extraFolderPaths);
@@ -2931,42 +3145,70 @@ export class SettingsManager {
} }
} }
editApiKey() { updateLlmApiKeyStatus() {
const statusEl = document.getElementById('civitaiApiKeyStatus'); const hasKey = !!(state.global.settings.llm_api_key_set || state.global.settings.llm_api_key);
const statusText = document.getElementById('llmApiKeyStatusText');
const actionBtn = document.getElementById('llmApiKeyActionBtn');
if (!statusText || !actionBtn) return;
if (hasKey) {
statusText.classList.remove('api-key-status--unconfigured');
statusText.classList.add('api-key-status--configured');
statusText.innerHTML = '<i class="fas fa-check-circle text-success"></i> '
+ translate('settings.aiProvider.apiKeyConfigured', {}, 'Configured');
actionBtn.textContent = translate('common.actions.change', {}, 'Change');
} else {
statusText.classList.remove('api-key-status--configured');
statusText.classList.add('api-key-status--unconfigured');
statusText.innerHTML = '<i class="fas fa-times-circle text-error"></i> '
+ translate('settings.aiProvider.apiKeyNotSet', {}, 'Not set');
actionBtn.textContent = translate('settings.aiProvider.apiKeySet', {}, 'Set up');
}
}
editApiKey(settingsKey = 'civitai_api_key', inputId = 'civitaiApiKey') {
const statusId = inputId + 'Status';
const editId = inputId + 'Edit';
const statusEl = document.getElementById(statusId);
if (statusEl) statusEl.classList.add('is-hidden'); if (statusEl) statusEl.classList.add('is-hidden');
const editContainer = document.getElementById('civitaiApiKeyEdit'); const editContainer = document.getElementById(editId);
if (editContainer) editContainer.classList.remove('is-hidden'); if (editContainer) editContainer.classList.remove('is-hidden');
// Focus the input // Focus the input
const input = document.getElementById('civitaiApiKey'); const input = document.getElementById(inputId);
if (input) { if (input) {
input.value = ''; // Never pre-fill the secret input.value = ''; // Never pre-fill the secret
setTimeout(() => input.focus(), 50); setTimeout(() => input.focus(), 50);
} }
} }
cancelEditApiKey(silent) { cancelEditApiKey(silent, inputId = 'civitaiApiKey') {
const editContainer = document.getElementById('civitaiApiKeyEdit'); const editId = inputId + 'Edit';
const statusId = inputId + 'Status';
const editContainer = document.getElementById(editId);
if (editContainer) editContainer.classList.add('is-hidden'); if (editContainer) editContainer.classList.add('is-hidden');
const statusContainer = document.getElementById('civitaiApiKeyStatus'); const statusContainer = document.getElementById(statusId);
if (statusContainer) statusContainer.classList.remove('is-hidden'); if (statusContainer) statusContainer.classList.remove('is-hidden');
// Clear any typed value // Clear any typed value
const input = document.getElementById('civitaiApiKey'); const input = document.getElementById(inputId);
if (input) input.value = ''; if (input) input.value = '';
if (!silent) { if (!silent) {
this.updateApiKeyStatus(); if (inputId === 'civitaiApiKey') {
this.updateApiKeyStatus();
}
} }
} }
async saveApiKey() { async saveApiKey(settingsKey = 'civitai_api_key', inputId = 'civitaiApiKey') {
const input = document.getElementById('civitaiApiKey'); const input = document.getElementById(inputId);
if (!input) return; if (!input) return;
const value = input.value.trim(); const value = input.value.trim();
try { try {
await this.saveSetting('civitai_api_key', value); await this.saveSetting(settingsKey, value);
const labelName = settingsKey === 'civitai_api_key' ? 'CivitAI API Key' : 'LLM API Key';
showToast('toast.settings.settingsUpdated', showToast('toast.settings.settingsUpdated',
{ setting: 'CivitAI API Key' }, 'success'); { setting: labelName }, 'success');
} catch (error) { } catch (error) {
showToast('toast.settings.settingSaveFailed', showToast('toast.settings.settingSaveFailed',
{ message: error.message }, 'error'); { message: error.message }, 'error');
@@ -2974,9 +3216,13 @@ export class SettingsManager {
} }
// Update the in-memory flag so the UI reflects the change // Update the in-memory flag so the UI reflects the change
state.global.settings.civitai_api_key_set = !!value; if (settingsKey === 'civitai_api_key') {
this.cancelEditApiKey(true); state.global.settings.civitai_api_key_set = !!value;
this.updateApiKeyStatus(); }
this.cancelEditApiKey(true, inputId);
if (inputId === 'civitaiApiKey') {
this.updateApiKeyStatus();
}
} }
toggleInputVisibility(button) { toggleInputVisibility(button) {
+29 -8
View File
@@ -168,6 +168,18 @@ export class DownloadManager {
let failedDownloads = 0; let failedDownloads = 0;
let accessFailures = 0; let accessFailures = 0;
let currentLoraProgress = 0; let currentLoraProgress = 0;
let cancelled = false;
this.importManager.loadingManager.showCancelButton(async () => {
if (cancelled) return;
cancelled = true;
try {
const loraClient = getModelApiClient(MODEL_TYPES.LORA);
await loraClient.cancelDownload(batchDownloadId);
} catch (e) {
console.error('Cancel request failed:', e);
}
});
// Set up progress tracking for current download // Set up progress tracking for current download
ws.onmessage = (event) => { ws.onmessage = (event) => {
@@ -179,6 +191,11 @@ export class DownloadManager {
return; return;
} }
if (data.status === 'cancelled') {
cancelled = true;
return;
}
// Process progress updates for our current active download // Process progress updates for our current active download
if (data.status === 'progress' && data.download_id && data.download_id.startsWith(batchDownloadId)) { if (data.status === 'progress' && data.download_id && data.download_id.startsWith(batchDownloadId)) {
// Update current LoRA progress // Update current LoRA progress
@@ -221,6 +238,8 @@ export class DownloadManager {
const useDefaultPaths = getStorageItem('use_default_path_loras', false); const useDefaultPaths = getStorageItem('use_default_path_loras', false);
for (let i = 0; i < this.importManager.downloadableLoRAs.length; i++) { for (let i = 0; i < this.importManager.downloadableLoRAs.length; i++) {
if (cancelled) break;
const lora = this.importManager.downloadableLoRAs[i]; const lora = this.importManager.downloadableLoRAs[i];
// Reset current LoRA progress for new download // Reset current LoRA progress for new download
@@ -241,15 +260,13 @@ export class DownloadManager {
batchDownloadId batchDownloadId
); );
if (cancelled) break;
if (!response.success) { if (!response.success) {
console.error(`Failed to download LoRA ${lora.name}: ${response.error}`); console.error(`Failed to download LoRA ${lora.name}: ${response.error}`);
failedDownloads++; failedDownloads++;
// Continue with next download
} else { } else {
completedDownloads++; completedDownloads++;
// Update progress to show completion of current LoRA
updateProgress(100, completedDownloads, ''); updateProgress(100, completedDownloads, '');
if (completedDownloads + failedDownloads < this.importManager.downloadableLoRAs.length) { if (completedDownloads + failedDownloads < this.importManager.downloadableLoRAs.length) {
@@ -259,9 +276,10 @@ export class DownloadManager {
} }
} }
} catch (downloadError) { } catch (downloadError) {
console.error(`Error downloading LoRA ${lora.name}:`, downloadError); if (!cancelled) {
failedDownloads++; console.error(`Error downloading LoRA ${lora.name}:`, downloadError);
// Continue with next download failedDownloads++;
}
} }
} }
@@ -269,7 +287,10 @@ export class DownloadManager {
ws.close(); ws.close();
// Show appropriate completion message based on results // Show appropriate completion message based on results
if (failedDownloads === 0) { if (cancelled) {
showToast('toast.downloads.downloadStopped', {}, 'info',
`Download cancelled. ${completedDownloads} item(s) completed.`);
} else if (failedDownloads === 0) {
showToast('toast.loras.allDownloadSuccessful', { count: completedDownloads }, 'success'); showToast('toast.loras.allDownloadSuccessful', { count: completedDownloads }, 'success');
} else { } else {
if (accessFailures > 0) { if (accessFailures > 0) {
+4
View File
@@ -55,6 +55,10 @@ const DEFAULT_SETTINGS_BASE = Object.freeze({
strip_lora_on_copy: false, strip_lora_on_copy: false,
use_new_license_icons: true, use_new_license_icons: true,
group_by_model: false, group_by_model: false,
llm_provider: 'openai',
llm_api_key: '',
llm_api_base: '',
llm_model: '',
}); });
export function createDefaultSettings() { export function createDefaultSettings() {
+34 -22
View File
@@ -66,11 +66,23 @@ export const BASE_MODELS = {
HUNYUAN_VIDEO: "Hunyuan Video", HUNYUAN_VIDEO: "Hunyuan Video",
// Other models // Other models
ANIMA: "Anima", ANIMA: "Anima",
ACE_AUDIO: "ACE Audio",
BOOGU: "Boogu",
ERNIE: "Ernie", ERNIE: "Ernie",
ERNIE_TURBO: "Ernie Turbo", ERNIE_TURBO: "Ernie Turbo",
NUCLEUS: "Nucleus", GROK: "Grok",
PONY_V7: "Pony V7", HAPPY_HORSE: "HappyHorse",
HIDREAM_O1: "HiDream-O1",
IDEOGRAM_4_0: "Ideogram 4.0",
KREA_2: "Krea 2", KREA_2: "Krea 2",
LENS: "Lens",
PONY_V7: "Pony V7",
MAI: "MAI",
NUCLEUS: "Nucleus",
QWEN_2: "Qwen 2",
UPSCALER: "Upscaler",
WAN_IMAGE_2_7: "Wan Image 2.7",
WAN_VIDEO_2_7: "Wan Video 2.7",
// Default // Default
UNKNOWN: "Other" UNKNOWN: "Other"
}; };
@@ -143,22 +155,6 @@ export const BASE_MODEL_ABBREVIATIONS = {
[BASE_MODELS.FLUX_2_KLEIN_4B]: 'FK4', [BASE_MODELS.FLUX_2_KLEIN_4B]: 'FK4',
[BASE_MODELS.FLUX_2_KLEIN_4B_BASE]: 'FK4B', [BASE_MODELS.FLUX_2_KLEIN_4B_BASE]: 'FK4B',
// Other diffusion models
[BASE_MODELS.AURAFLOW]: 'AF',
[BASE_MODELS.CHROMA]: 'CHR',
[BASE_MODELS.PIXART_A]: 'PXA',
[BASE_MODELS.PIXART_E]: 'PXE',
[BASE_MODELS.HUNYUAN_1]: 'HY',
[BASE_MODELS.LUMINA]: 'L',
[BASE_MODELS.KOLORS]: 'KLR',
[BASE_MODELS.NOOBAI]: 'NAI',
[BASE_MODELS.ILLUSTRIOUS]: 'IL',
[BASE_MODELS.PONY]: 'PONY',
[BASE_MODELS.HIDREAM]: 'HID',
[BASE_MODELS.QWEN]: 'QWEN',
[BASE_MODELS.ZIMAGE_TURBO]: 'ZIT',
[BASE_MODELS.ZIMAGE_BASE]: 'ZIB',
// Video models // Video models
[BASE_MODELS.SVD]: 'SVD', [BASE_MODELS.SVD]: 'SVD',
[BASE_MODELS.LTXV]: 'LTXV', [BASE_MODELS.LTXV]: 'LTXV',
@@ -195,10 +191,22 @@ export const BASE_MODEL_ABBREVIATIONS = {
[BASE_MODELS.ZIMAGE_TURBO]: 'ZIT', [BASE_MODELS.ZIMAGE_TURBO]: 'ZIT',
[BASE_MODELS.ZIMAGE_BASE]: 'ZIB', [BASE_MODELS.ZIMAGE_BASE]: 'ZIB',
[BASE_MODELS.ANIMA]: 'ANI', [BASE_MODELS.ANIMA]: 'ANI',
[BASE_MODELS.ACE_AUDIO]: 'ACE',
[BASE_MODELS.BOOGU]: 'BOOG',
[BASE_MODELS.ERNIE]: 'ERNI', [BASE_MODELS.ERNIE]: 'ERNI',
[BASE_MODELS.ERNIE_TURBO]: 'ETRB', [BASE_MODELS.ERNIE_TURBO]: 'ETRB',
[BASE_MODELS.NUCLEUS]: 'NUCL', [BASE_MODELS.GROK]: 'GROK',
[BASE_MODELS.HAPPY_HORSE]: 'HAPP',
[BASE_MODELS.HIDREAM_O1]: 'HIO1',
[BASE_MODELS.IDEOGRAM_4_0]: 'ID40',
[BASE_MODELS.KREA_2]: 'KR2', [BASE_MODELS.KREA_2]: 'KR2',
[BASE_MODELS.LENS]: 'LENS',
[BASE_MODELS.MAI]: 'MAI',
[BASE_MODELS.NUCLEUS]: 'NUCL',
[BASE_MODELS.QWEN_2]: 'QWN2',
[BASE_MODELS.UPSCALER]: 'UPSC',
[BASE_MODELS.WAN_IMAGE_2_7]: 'WI27',
[BASE_MODELS.WAN_VIDEO_2_7]: 'WAN',
// Default // Default
[BASE_MODELS.UNKNOWN]: 'OTH' [BASE_MODELS.UNKNOWN]: 'OTH'
@@ -394,7 +402,9 @@ export const BASE_MODEL_CATEGORIES = {
BASE_MODELS.WAN_VIDEO_14B_I2V_480P, BASE_MODELS.WAN_VIDEO_14B_I2V_720P, BASE_MODELS.WAN_VIDEO_14B_I2V_480P, BASE_MODELS.WAN_VIDEO_14B_I2V_720P,
BASE_MODELS.WAN_VIDEO_2_2_TI2V_5B, BASE_MODELS.WAN_VIDEO_2_2_T2V_A14B, BASE_MODELS.WAN_VIDEO_2_2_TI2V_5B, BASE_MODELS.WAN_VIDEO_2_2_T2V_A14B,
BASE_MODELS.WAN_VIDEO_2_2_I2V_A14B, BASE_MODELS.WAN_VIDEO_2_5_T2V, BASE_MODELS.WAN_VIDEO_2_2_I2V_A14B, BASE_MODELS.WAN_VIDEO_2_5_T2V,
BASE_MODELS.WAN_VIDEO_2_5_I2V BASE_MODELS.WAN_VIDEO_2_5_I2V,
BASE_MODELS.HAPPY_HORSE,
BASE_MODELS.WAN_IMAGE_2_7, BASE_MODELS.WAN_VIDEO_2_7
], ],
'Flux Models': [BASE_MODELS.FLUX_1_D, BASE_MODELS.FLUX_1_S, BASE_MODELS.FLUX_1_KONTEXT, BASE_MODELS.FLUX_1_KREA, BASE_MODELS.FLUX_2_D, BASE_MODELS.FLUX_2_KLEIN_9B, BASE_MODELS.FLUX_2_KLEIN_9B_BASE, BASE_MODELS.FLUX_2_KLEIN_4B, BASE_MODELS.FLUX_2_KLEIN_4B_BASE], 'Flux Models': [BASE_MODELS.FLUX_1_D, BASE_MODELS.FLUX_1_S, BASE_MODELS.FLUX_1_KONTEXT, BASE_MODELS.FLUX_1_KREA, BASE_MODELS.FLUX_2_D, BASE_MODELS.FLUX_2_KLEIN_9B, BASE_MODELS.FLUX_2_KLEIN_9B_BASE, BASE_MODELS.FLUX_2_KLEIN_4B, BASE_MODELS.FLUX_2_KLEIN_4B_BASE],
'Other Models': [ 'Other Models': [
@@ -402,8 +412,10 @@ export const BASE_MODEL_CATEGORIES = {
BASE_MODELS.QWEN, BASE_MODELS.AURAFLOW, BASE_MODELS.CHROMA, BASE_MODELS.ZIMAGE_TURBO, BASE_MODELS.ZIMAGE_BASE, BASE_MODELS.QWEN, BASE_MODELS.AURAFLOW, BASE_MODELS.CHROMA, BASE_MODELS.ZIMAGE_TURBO, BASE_MODELS.ZIMAGE_BASE,
BASE_MODELS.PIXART_A, BASE_MODELS.PIXART_E, BASE_MODELS.HUNYUAN_1, BASE_MODELS.PIXART_A, BASE_MODELS.PIXART_E, BASE_MODELS.HUNYUAN_1,
BASE_MODELS.LUMINA, BASE_MODELS.KOLORS, BASE_MODELS.NOOBAI, BASE_MODELS.ANIMA, BASE_MODELS.LUMINA, BASE_MODELS.KOLORS, BASE_MODELS.NOOBAI, BASE_MODELS.ANIMA,
BASE_MODELS.ERNIE, BASE_MODELS.ERNIE_TURBO, BASE_MODELS.NUCLEUS, BASE_MODELS.ACE_AUDIO, BASE_MODELS.BOOGU, BASE_MODELS.ERNIE, BASE_MODELS.ERNIE_TURBO,
BASE_MODELS.KREA_2, BASE_MODELS.GROK, BASE_MODELS.HIDREAM_O1, BASE_MODELS.IDEOGRAM_4_0,
BASE_MODELS.LENS, BASE_MODELS.MAI, BASE_MODELS.NUCLEUS,
BASE_MODELS.QWEN_2, BASE_MODELS.KREA_2, BASE_MODELS.UPSCALER,
BASE_MODELS.UNKNOWN BASE_MODELS.UNKNOWN
] ]
}; };
+39
View File
@@ -552,6 +552,8 @@ async function fetchWorkflowRegistry() {
if (!registryData.success) { if (!registryData.success) {
if (registryData.error === 'Standalone Mode Active') { if (registryData.error === 'Standalone Mode Active') {
showToast('toast.general.cannotInteractStandalone', {}, 'warning'); showToast('toast.general.cannotInteractStandalone', {}, 'warning');
} else if (registryData.error === 'Empty Registry') {
showToast('uiHelpers.workflow.noSupportedNodes', {}, 'warning');
} else { } else {
showToast('toast.general.failedWorkflowInfo', {}, 'error'); showToast('toast.general.failedWorkflowInfo', {}, 'error');
} }
@@ -1482,3 +1484,40 @@ export async function openExampleImagesFolder(modelHash) {
return false; return false;
} }
} }
/**
* Set up a paste handler on a textarea that automatically appends a newline
* after pasted content that looks like a URL (http/https). This lets users
* paste multiple URLs one after another without manually pressing Enter.
* @param {string} textareaId - The id of the textarea element
*/
export function setupAutoNewlineOnPaste(textareaId) {
const el = document.getElementById(textareaId);
if (!el || el.tagName !== 'TEXTAREA') return;
el.addEventListener('paste', (e) => {
const pastedText = (e.clipboardData || window.clipboardData).getData('text');
// Only apply to text that starts with http:// or https://
if (/^https?:\/\//.test(pastedText) && !pastedText.endsWith('\n')) {
e.preventDefault();
const start = el.selectionStart;
const end = el.selectionEnd;
const text = el.value;
const before = text.substring(0, start);
const after = text.substring(end);
// Append newline after the pasted URL
const modifiedText = pastedText + '\n';
el.value = before + modifiedText + after;
// Move cursor to just after the inserted text
const newCursorPos = start + modifiedText.length;
el.selectionStart = el.selectionEnd = newCursorPos;
// Trigger input event so any listeners stay in sync
el.dispatchEvent(new Event('input', { bubbles: true }));
}
// Non-URL text or text already ending with \n — let default paste happen
});
}
+13 -1
View File
@@ -12,7 +12,19 @@
<div id="checkpointContextMenu" class="context-menu" style="display: none;"> <div id="checkpointContextMenu" class="context-menu" style="display: none;">
<!-- Metadata --> <!-- Metadata -->
<div class="context-menu-item" data-action="refresh-metadata"><i class="fas fa-sync"></i> {{ t('loras.contextMenu.refreshMetadata') }}</div> <div class="context-menu-item" data-action="refresh-metadata"><i class="fas fa-sync"></i> {{ t('loras.contextMenu.refreshMetadata') }}</div>
<div class="context-menu-item" data-action="relink-civitai"><i class="fas fa-link"></i> {{ t('loras.contextMenu.relinkCivitai') }}</div> <div class="context-menu-item has-submenu" data-has-submenu="link-model">
<i class="fas fa-link"></i>
<span>{{ t('loras.contextMenu.linkModel') }}</span>
<i class="fas fa-chevron-right submenu-arrow"></i>
<div class="context-submenu">
<div class="context-menu-item" data-action="relink-civitai">
<i class="fas fa-external-link-alt"></i> <span>{{ t('loras.contextMenu.linkCivitai') }}</span>
</div>
<div class="context-menu-item" data-action="link-hf">
<i class="fas fa-robot"></i> <span>{{ t('loras.contextMenu.linkHuggingFace') }}</span>
</div>
</div>
</div>
<div class="context-menu-separator menu-section-break"></div> <div class="context-menu-separator menu-section-break"></div>
<!-- Workflow --> <!-- Workflow -->
<div class="context-menu-item" data-action="copyname"><i class="fas fa-copy"></i> {{ t('loras.contextMenu.copyFilename') }}</div> <div class="context-menu-item" data-action="copyname"><i class="fas fa-copy"></i> {{ t('loras.contextMenu.copyFilename') }}</div>
+18 -2
View File
@@ -12,8 +12,21 @@
<div class="context-menu-item" data-action="check-updates"> <div class="context-menu-item" data-action="check-updates">
<i class="fas fa-bell"></i> <span>{{ t('loras.contextMenu.checkUpdates') }}</span> <i class="fas fa-bell"></i> <span>{{ t('loras.contextMenu.checkUpdates') }}</span>
</div> </div>
<div class="context-menu-item" data-action="relink-civitai"> <div class="context-menu-item has-submenu" data-has-submenu="link-model">
<i class="fas fa-link"></i> <span>{{ t('loras.contextMenu.relinkCivitai') }}</span> <i class="fas fa-link"></i>
<span>{{ t('loras.contextMenu.linkModel') }}</span>
<i class="fas fa-chevron-right submenu-arrow"></i>
<div class="context-submenu">
<div class="context-menu-item" data-action="relink-civitai">
<i class="fas fa-external-link-alt"></i> <span>{{ t('loras.contextMenu.linkCivitai') }}</span>
</div>
<div class="context-menu-item" data-action="link-hf">
<i class="fas fa-robot"></i> <span>{{ t('loras.contextMenu.linkHuggingFace') }}</span>
</div>
</div>
</div>
<div class="context-menu-item" data-action="enrich-hf-llm">
<i class="fas fa-wand-magic-sparkles"></i> <span>{{ t('loras.contextMenu.enrichHfAgent') }}</span>
</div> </div>
<div class="context-menu-separator menu-section-break"></div> <div class="context-menu-separator menu-section-break"></div>
<!-- Workflow --> <!-- Workflow -->
@@ -83,6 +96,9 @@
<div class="context-menu-item" data-action="resume-metadata-refresh"> <div class="context-menu-item" data-action="resume-metadata-refresh">
<i class="fas fa-redo"></i> <span>{{ t('loras.bulkOperations.resumeMetadataRefresh') }}</span> <i class="fas fa-redo"></i> <span>{{ t('loras.bulkOperations.resumeMetadataRefresh') }}</span>
</div> </div>
<div class="context-menu-item" data-action="enrich-hf-llm-bulk">
<i class="fas fa-wand-magic-sparkles"></i> <span>{{ t('loras.bulkOperations.enrichHfAgent') }}</span>
</div>
</div> </div>
<div class="context-menu-section" data-section="workflow"> <div class="context-menu-section" data-section="workflow">
<div class="context-menu-section-header">{{ t('loras.bulkOperations.sections.workflow') }}</div> <div class="context-menu-section-header">{{ t('loras.bulkOperations.sections.workflow') }}</div>
+1
View File
@@ -8,6 +8,7 @@
{% include 'components/modals/update_modal.html' %} {% include 'components/modals/update_modal.html' %}
{% include 'components/modals/help_modal.html' %} {% include 'components/modals/help_modal.html' %}
{% include 'components/modals/relink_civitai_modal.html' %} {% include 'components/modals/relink_civitai_modal.html' %}
{% include 'components/modals/link_hf_modal.html' %}
{% include 'components/modals/example_access_modal.html' %} {% include 'components/modals/example_access_modal.html' %}
{% include 'components/modals/download_modal.html' %} {% include 'components/modals/download_modal.html' %}
{% include 'components/modals/move_modal.html' %} {% include 'components/modals/move_modal.html' %}
@@ -112,6 +112,10 @@
<a href="https://github.com/willmiao/ComfyUI-Lora-Manager/wiki/Priority-Tags-Configuration-Guide" target="_blank"> <a href="https://github.com/willmiao/ComfyUI-Lora-Manager/wiki/Priority-Tags-Configuration-Guide" target="_blank">
Priority Tags Configuration Guide Priority Tags Configuration Guide
<span class="new-content-badge inline">{{ t('help.documentation.newBadge') }}</span> <span class="new-content-badge inline">{{ t('help.documentation.newBadge') }}</span>
<li>
<a href="https://github.com/willmiao/ComfyUI-Lora-Manager/wiki/AI-Provider-Setup" target="_blank">
AI Provider Setup
<span class="new-content-badge inline">{{ t('help.documentation.newBadge') }}</span>
</a> </a>
</li> </li>
</ul> </ul>
@@ -0,0 +1,24 @@
<!-- Link to HuggingFace Modal -->
<div id="linkHfModal" class="modal">
<div class="modal-content">
<button class="close" onclick="modalManager.closeModal('linkHfModal')">&times;</button>
<h2>{{ t('modals.linkHuggingFace.title') }}</h2>
<div class="warning-box">
<i class="fas fa-info-circle"></i>
<p>{{ t('modals.linkHuggingFace.infoText') }}</p>
</div>
<div class="input-group">
<label for="hfModelUrl">{{ t('modals.linkHuggingFace.urlLabel') }}</label>
<input type="text" id="hfModelUrl" placeholder="{{ t('modals.linkHuggingFace.urlPlaceholder') }}" />
<div class="input-error" id="hfModelUrlError"></div>
<div class="input-help">
{{ t('modals.linkHuggingFace.helpText') }}<br>
<strong>https://huggingface.co/user/repo</strong>
</div>
</div>
<div class="modal-actions">
<button class="cancel-btn" onclick="modalManager.closeModal('linkHfModal')">{{ t('common.actions.cancel') }}</button>
<button class="confirm-btn" id="confirmLinkHfBtn">{{ t('modals.linkHuggingFace.confirmAction') }}</button>
</div>
</div>
</div>
@@ -144,6 +144,112 @@
</div> </div>
</div> </div>
<!-- AI Provider Configuration (BYOK) -->
<div class="settings-subsection">
<div class="settings-subsection-header">
<h4>{{ t('settings.aiProvider.title') }}</h4>
</div>
<div class="setting-item">
<div class="setting-row">
<div class="setting-info">
<label for="llmProvider">{{ t('settings.aiProvider.provider') }}</label>
<i class="fas fa-info-circle info-icon" data-tooltip="{{ t('settings.aiProvider.providerHelp') }}"></i>
</div>
<div class="setting-control select-control">
<select id="llmProvider" onchange="settingsManager.saveSelectSetting('llmProvider', 'llm_provider')">
<option value="openai">{{ t('settings.aiProvider.providerOptions.openai') }}</option>
<option value="ollama">{{ t('settings.aiProvider.providerOptions.ollama') }}</option>
<option value="deepseek">{{ t('settings.aiProvider.providerOptions.deepseek') }}</option>
<option value="groq">{{ t('settings.aiProvider.providerOptions.groq') }}</option>
<option value="openrouter">{{ t('settings.aiProvider.providerOptions.openrouter') }}</option>
<option value="opencode-go">{{ t('settings.aiProvider.providerOptions.opencode-go') }}</option>
<option value="custom">{{ t('settings.aiProvider.providerOptions.custom') }}</option>
</select>
</div>
</div>
</div>
<div class="setting-item">
<div class="setting-row">
<div class="setting-info">
<label for="llmApiBase">{{ t('settings.aiProvider.apiBase') }}</label>
<i class="fas fa-info-circle info-icon" data-tooltip="{{ t('settings.aiProvider.apiBaseHelp') }}"></i>
</div>
<div class="setting-control">
<div class="text-input-wrapper lm-combobox-container">
<input type="text" id="llmApiBase"
class="lm-combobox-input"
value="{{ settings.get('llm_api_base', '') }}"
placeholder="{{ t('settings.aiProvider.apiBasePlaceholder') }}"
autocomplete="off"
onblur="settingsManager.saveInputSetting('llmApiBase', 'llm_api_base')"
onkeydown="if(event.key === 'Enter') { this.blur(); }" />
</div>
</div>
</div>
</div>
<div class="setting-item api-key-item">
<div class="setting-row">
<div class="setting-info">
<label>{{ t('settings.aiProvider.apiKey') }}</label>
<i class="fas fa-info-circle info-icon" data-tooltip="{{ t('settings.aiProvider.apiKeyHelp') }}"></i>
</div>
<div class="setting-control">
<div id="llmApiKeyStatus" class="api-key-status">
<span id="llmApiKeyStatusText" class="api-key-status-text api-key-status--unconfigured">
<i class="fas fa-times-circle text-error"></i>
{{ t('settings.aiProvider.apiKeyNotSet') }}
</span>
<button type="button" class="secondary-btn" id="llmApiKeyActionBtn" onclick="settingsManager.editApiKey('llm_api_key', 'llmApiKey')">
{{ t('settings.aiProvider.apiKeySet') }}
</button>
</div>
<div id="llmApiKeyEdit" class="api-key-edit is-hidden">
<div class="api-key-input">
<input type="text"
id="llmApiKey"
class="api-key-masked"
placeholder="{{ t('settings.aiProvider.apiKeyPlaceholder') }}"
autocomplete="off"
data-mask="css" />
<button type="button" class="toggle-visibility">
<i class="fas fa-eye"></i>
</button>
</div>
<button type="button" class="primary-btn" onclick="settingsManager.saveApiKey('llm_api_key', 'llmApiKey')">{{ t('common.actions.save') }}</button>
<button type="button" class="secondary-btn" onclick="settingsManager.cancelEditApiKey(true, 'llmApiKey')">{{ t('common.actions.cancel') }}</button>
</div>
</div>
</div>
</div>
<div class="setting-item">
<div class="setting-row">
<div class="setting-info">
<label for="llmModel">{{ t('settings.aiProvider.model') }}</label>
<i class="fas fa-info-circle info-icon" data-tooltip="{{ t('settings.aiProvider.modelHelp') }}"></i>
</div>
<div class="setting-control">
<div class="text-input-wrapper lm-combobox-container">
<input type="text" id="llmModel"
class="lm-combobox-input"
value="{{ settings.get('llm_model', '') }}"
placeholder="{{ t('settings.aiProvider.modelPlaceholder') }}"
autocomplete="off"
onblur="settingsManager.saveInputSetting('llmModel', 'llm_model')"
onkeydown="if(event.key === 'Enter') { this.blur(); }" />
</div>
</div>
</div>
</div>
</div>
<!-- Provider presets + model lists for frontend -->
<script id="llmProviderPresets" type="application/json">
{{ provider_presets_json | safe }}
</script>
<script id="llmProviderModels" type="application/json">
{{ provider_models_json | safe }}
</script>
<div class="settings-subsection"> <div class="settings-subsection">
<div class="settings-subsection-header"> <div class="settings-subsection-header">
<h4>{{ t('settings.sections.downloads') }}</h4> <h4>{{ t('settings.sections.downloads') }}</h4>
+13 -1
View File
@@ -12,7 +12,19 @@
<div id="embeddingContextMenu" class="context-menu" style="display: none;"> <div id="embeddingContextMenu" class="context-menu" style="display: none;">
<!-- Metadata --> <!-- Metadata -->
<div class="context-menu-item" data-action="refresh-metadata"><i class="fas fa-sync"></i> {{ t('loras.contextMenu.refreshMetadata') }}</div> <div class="context-menu-item" data-action="refresh-metadata"><i class="fas fa-sync"></i> {{ t('loras.contextMenu.refreshMetadata') }}</div>
<div class="context-menu-item" data-action="relink-civitai"><i class="fas fa-link"></i> {{ t('loras.contextMenu.relinkCivitai') }}</div> <div class="context-menu-item has-submenu" data-has-submenu="link-model">
<i class="fas fa-link"></i>
<span>{{ t('loras.contextMenu.linkModel') }}</span>
<i class="fas fa-chevron-right submenu-arrow"></i>
<div class="context-submenu">
<div class="context-menu-item" data-action="relink-civitai">
<i class="fas fa-external-link-alt"></i> <span>{{ t('loras.contextMenu.linkCivitai') }}</span>
</div>
<div class="context-menu-item" data-action="link-hf">
<i class="fas fa-robot"></i> <span>{{ t('loras.contextMenu.linkHuggingFace') }}</span>
</div>
</div>
</div>
<div class="context-menu-separator menu-section-break"></div> <div class="context-menu-separator menu-section-break"></div>
<!-- Workflow --> <!-- Workflow -->
<div class="context-menu-item" data-action="copyname"><i class="fas fa-copy"></i> {{ t('loras.contextMenu.copyFilename') }}</div> <div class="context-menu-item" data-action="copyname"><i class="fas fa-copy"></i> {{ t('loras.contextMenu.copyFilename') }}</div>
+1
View File
@@ -0,0 +1 @@
# Test suite package.
+1
View File
@@ -0,0 +1 @@
# HF Metadata Enrichment validation suite.
+133
View File
@@ -0,0 +1,133 @@
"""Configuration for the HF metadata enrichment validation suite.
Loads user settings, defines paths, and pulls constants from the main
codebase (``py.utils.constants``).
"""
from __future__ import annotations
import json
import logging
import os
from typing import Any, Dict, List
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Default paths
# ---------------------------------------------------------------------------
_DEFAULT_MODELS_FILE = os.path.join(
os.path.dirname(__file__), "test_data", "hf_lora_models_with_safetensors.txt"
)
_DEFAULT_SETTINGS_PATH = os.path.expanduser(
"~/.config/ComfyUI-LoRA-Manager/settings.json"
)
_DEFAULT_OUTPUT_DIR = "/tmp/hf_enrich_validation"
# ---------------------------------------------------------------------------
# Constants from the main codebase (copied at import time)
# ---------------------------------------------------------------------------
# Priority tags used in the LLM prompt for tag selection guidance.
CIVITAI_MODEL_TAGS: List[str] = [
"character", "concept", "clothing", "realistic", "anime", "toon",
"furry", "style", "poses", "background", "tool", "vehicle",
"buildings", "objects", "assets", "animal", "action",
]
# ---------------------------------------------------------------------------
# Base model resolution — dynamically fetched from production code
# ---------------------------------------------------------------------------
# Module-level cache — populated by init_supported_base_models().
# Falls back to a comprehensive hardcoded list when the live fetch fails.
SUPPORTED_BASE_MODELS: List[str] = []
# Fallback base models when the production list_base_models() is unavailable.
_FALLBACK_BASE_MODELS: List[str] = [
"SD 1.4", "SD 1.5", "SD 1.5 LCM", "SD 1.5 Hyper",
"SD 2.0", "SD 2.1",
"SD 3", "SD 3.5", "SD 3.5 Medium", "SD 3.5 Large", "SD 3.5 Large Turbo",
"SDXL 1.0", "SDXL Lightning", "SDXL Hyper",
"Flux.1 D", "Flux.1 S", "Flux.1 Krea", "Flux.1 Kontext",
"Flux.2 D", "Flux.2 Klein 9B", "Flux.2 Klein 9B-base",
"Flux.2 Klein 4B", "Flux.2 Klein 4B-base",
"AuraFlow", "Chroma", "PixArt a", "PixArt E",
"Hunyuan 1", "Lumina", "Kolors",
"NoobAI", "Illustrious", "Pony", "Pony V7",
"HiDream", "Qwen", "ZImageTurbo", "ZImageBase",
"SVD", "LTXV", "LTXV2", "LTXV 2.3",
"CogVideoX", "Mochi",
"Wan Video", "Wan Video 1.3B t2v", "Wan Video 14B t2v",
"Wan Video 14B i2v 480p", "Wan Video 14B i2v 720p",
"Wan Video 2.2 TI2V-5B", "Wan Video 2.2 T2V-A14B",
"Wan Video 2.2 I2V-A14B",
"Wan Video 2.5 T2V", "Wan Video 2.5 I2V",
"Hunyuan Video", "Anima", "Ernie", "Ernie Turbo",
"Nucleus", "Krea 2",
]
async def init_supported_base_models() -> None:
"""Populate ``SUPPORTED_BASE_MODELS`` from the production codebase.
Calls ``py.metadata_ops.list_base_models()`` which merges a hardcoded
fallback with models fetched from the CivitAI API. When the call
fails (e.g. offline, API error), falls back to ``_FALLBACK_BASE_MODELS``.
Must be called from within an async event loop (i.e. during
``run_validation.main()``, not at module level).
"""
try:
from py.metadata_ops import list_base_models
models = await list_base_models()
if models:
SUPPORTED_BASE_MODELS[:] = models
logger.info("Loaded %d base models from production code", len(models))
return
logger.warning("list_base_models returned empty list, using fallback")
except Exception as exc:
logger.warning("Failed to load base models from production: %s", exc)
SUPPORTED_BASE_MODELS[:] = _FALLBACK_BASE_MODELS
logger.info("Using fallback base model list (%d entries)", len(SUPPORTED_BASE_MODELS))
# Placeholder values the LLM sometimes emits that should count as "empty".
PLACEHOLDER_VALUES = frozenset({
"none", "null", "n/a", "unknown", "not available",
"not specified", "no trigger words", "no trigger word",
})
# ---------------------------------------------------------------------------
# User settings loader
# ---------------------------------------------------------------------------
def load_settings(settings_path: str) -> Dict[str, Any]:
"""Load LoRA Manager settings from *settings_path*.
Returns a flat dict with the LLM configuration fields that the
enrichment pipeline depends on.
"""
path = os.path.expanduser(settings_path)
if not os.path.exists(path):
raise FileNotFoundError(
f"Settings file not found: {path}\n"
"Please provide a valid --settings path."
)
with open(path, "r", encoding="utf-8") as fh:
raw: Dict[str, Any] = json.load(fh)
# Extract LLM-relevant config
return {
"llm_provider": raw.get("llm_provider", "ollama"),
"llm_model": raw.get("llm_model", "qwen3.5:9b"),
"llm_api_base": raw.get("llm_api_base", "http://localhost:11434/v1"),
"llm_api_key": raw.get("llm_api_key", ""),
"settings_path": path,
}
@@ -0,0 +1,208 @@
"""Execute the ``enrich_hf_metadata`` skill serially over a list of models.
Design decisions (local Ollama, no rate limits):
- Sequential execution: one model at a time. 100 models at ~30-90 s/call
roughly 1-2 h total.
- Progress persisted to a JSON checkpoint file so the run can be resumed
with ``--resume``.
- Per-model timeout guards against a stuck Ollama inference.
"""
from __future__ import annotations
import asyncio
import json
import logging
import os
import time
from typing import Any, Dict, List, Optional
logger = logging.getLogger(__name__)
_SKILL_NAME = "enrich_hf_metadata"
# How long to wait for a single LLM call before marking it timed-out.
_PER_MODEL_TIMEOUT = 240 # seconds
# ---------------------------------------------------------------------------
# Progress checkpoint helpers
# ---------------------------------------------------------------------------
_PROGRESS_FILE = "progress.json"
def _load_progress(output_dir: str) -> Dict[str, Any]:
path = os.path.join(output_dir, _PROGRESS_FILE)
if os.path.exists(path):
with open(path, "r") as fh:
return json.load(fh)
return {"completed": [], "failed": [], "timed_out": []}
def _save_progress(output_dir: str, progress: Dict[str, Any]) -> None:
path = os.path.join(output_dir, _PROGRESS_FILE)
with open(path, "w") as fh:
json.dump(progress, fh, indent=2)
# ---------------------------------------------------------------------------
# Core runner
# ---------------------------------------------------------------------------
class EnrichmentRunner:
"""Serial enrichment runner with checkpoint resume."""
def __init__(
self,
output_dir: str,
*,
per_model_timeout: int = _PER_MODEL_TIMEOUT,
) -> None:
self._output_dir = output_dir
self._per_model_timeout = per_model_timeout
self._agent_service: Optional[Any] = None
async def _ensure_agent_service(self) -> Any:
"""Lazy-init AgentService (expensive — needs LLMService init)."""
if self._agent_service is not None:
return self._agent_service
from py.services.agent.agent_service import AgentService
self._agent_service = await AgentService.get_instance()
return self._agent_service
async def run(
self,
model_paths: List[str],
repos: List[str],
) -> Dict[str, Any]:
"""Run enrichment over *model_paths* (one-by-one).
Args:
model_paths: model paths in the same order as *repos*.
repos: HF repo IDs (for display / checkpoint labelling).
Returns:
A dict with keys ``results``, ``progress``, ``durations``.
"""
assert len(model_paths) == len(repos)
progress = _load_progress(self._output_dir)
completed_set = set(progress["completed"])
failed_set = set(progress["failed"])
timed_out_set = set(progress.get("timed_out", []))
agent = await self._ensure_agent_service()
results: List[Dict[str, Any]] = []
durations: Dict[str, float] = {}
total = len(model_paths)
processed_before = len(completed_set | failed_set | timed_out_set)
logger.info(
"Enrichment runner: %d models total, %d already processed",
total,
processed_before,
)
for idx, (model_path, repo_id) in enumerate(zip(model_paths, repos)):
if repo_id in completed_set:
logger.info("[%d/%d] SKIP (already done): %s", idx + 1, total, repo_id)
continue
if repo_id in failed_set or repo_id in timed_out_set:
logger.info(
"[%d/%d] SKIP (previously failed/timeout): %s",
idx + 1, total, repo_id,
)
continue
logger.info(
"[%d/%d] Enriching %s ...", idx + 1, total, repo_id,
)
t0 = time.perf_counter()
try:
result = await asyncio.wait_for(
agent.execute_skill(
skill_name=_SKILL_NAME,
input_data={"model_paths": [model_path]},
progress_callback=None,
),
timeout=self._per_model_timeout,
)
elapsed = time.perf_counter() - t0
durations[repo_id] = round(elapsed, 2)
if result.success:
completed_set.add(repo_id)
progress["completed"].append(repo_id)
logger.info(
"%s (%.1f s) — %s",
repo_id, elapsed, result.summary,
)
else:
failed_set.add(repo_id)
progress["failed"].append(repo_id)
logger.warning(
"%s (%.1f s) — %s",
repo_id, elapsed,
"; ".join(result.errors) if result.errors else result.summary,
)
results.append({
"repo_id": repo_id,
"model_path": model_path,
"success": result.success,
"updated_fields": result.updated_models,
"errors": result.errors,
"summary": result.summary,
"duration_s": round(elapsed, 2),
})
except asyncio.TimeoutError:
elapsed = time.perf_counter() - t0
durations[repo_id] = round(elapsed, 2)
timed_out_set.add(repo_id)
progress.setdefault("timed_out", []).append(repo_id)
logger.warning(
" ⏱ TIMEOUT %s (%.1f s, limit=%ds)",
repo_id, elapsed, self._per_model_timeout,
)
results.append({
"repo_id": repo_id,
"model_path": model_path,
"success": False,
"errors": [f"Timeout after {self._per_model_timeout}s"],
"summary": "LLM call timed out",
"duration_s": round(elapsed, 2),
})
except Exception as exc:
elapsed = time.perf_counter() - t0
durations[repo_id] = round(elapsed, 2)
failed_set.add(repo_id)
progress["failed"].append(repo_id)
logger.error(
"%s (%.1f s) — %s",
repo_id, elapsed, exc,
)
results.append({
"repo_id": repo_id,
"model_path": model_path,
"success": False,
"errors": [str(exc)],
"summary": f"Exception: {exc}",
"duration_s": round(elapsed, 2),
})
# Checkpoint after each model
_save_progress(self._output_dir, progress)
return {
"results": results,
"progress": progress,
"durations": durations,
}
@@ -0,0 +1,352 @@
"""Evaluate enriched ``.metadata.json`` quality across multiple dimensions.
Scoring rubric (per field):
- **Completeness**: Is the field populated with meaningful content?
- **Validity**: Does the value conform to expected constraints (controlled
vocab, non-placeholder, parsable JSON)?
- **Accuracy**: (sub-sample only requires manual verification against
the HF README).
"""
from __future__ import annotations
import json
import logging
import os
from typing import Any, Dict, List, Optional, Set
from .config import (
CIVITAI_MODEL_TAGS,
PLACEHOLDER_VALUES,
SUPPORTED_BASE_MODELS,
)
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Scoring helpers
# ---------------------------------------------------------------------------
_MIN_TAGS = 1
_MAX_TAGS = 8
_MIN_DESC_LENGTH = 20
_MIN_NOTES_LENGTH = 30
# Tags that the LLM sometimes emits but which are not meaningful content tags.
_TECH_TAGS = frozenset({
"lora", "dreambooth", "text-to-image", "diffusers", "flux",
"sdxl", "checkpoint", "pytorch", "safetensors", "fine-tuning",
"stable-diffusion", "training", "stablediffusion",
})
def _is_placeholder(val: str) -> bool:
return val.strip().lower() in PLACEHOLDER_VALUES
def _is_valid_trigger_words(words: List[str]) -> bool:
"""Return True if *words* is a non-empty list of real trigger words."""
if not words:
return False
cleaned = [w.strip() for w in words if w.strip()]
if not cleaned:
return False
# Reject if ALL entries are placeholders
non_placeholder = [w for w in cleaned if not _is_placeholder(w)]
return len(non_placeholder) > 0
def _is_valid_tags(tags: List[str]) -> bool:
"""Return True if *tags* is a reasonable list of content tags."""
if not tags:
return False
cleaned = [t.strip().lower() for t in tags if t.strip()]
if not cleaned:
return False
# At least one tag that isn't a technical keyword
meaningful = [t for t in cleaned if t not in _TECH_TAGS]
return len(meaningful) >= _MIN_TAGS
def _tag_priority_coverage(tags: List[str]) -> float:
"""Fraction of tags that align with the user's priority tag vocabulary."""
if not tags:
return 0.0
priority_lower = {t.lower() for t in CIVITAI_MODEL_TAGS}
matched = sum(1 for t in tags if t.strip().lower() in priority_lower)
return matched / len(tags)
# ---------------------------------------------------------------------------
# Per-model evaluation
# ---------------------------------------------------------------------------
# Type alias for a score record
ScoreRecord = Dict[str, Any]
def evaluate_model(
metadata: Dict[str, Any],
model_path: str,
repo_id: str,
*,
enrichment_success: bool,
enrichment_errors: List[str],
) -> ScoreRecord:
"""Score a single enriched model's metadata.
Returns a dict with per-field scores, a total score, and a list of
flagged issues.
"""
civitai = metadata.get("civitai") or {}
trained_words: List[str] = civitai.get("trainedWords") or metadata.get("trainedWords") or []
short_desc: str = civitai.get("description") or ""
tags: List[str] = metadata.get("tags") or []
notes: str = metadata.get("notes") or ""
usage_tips_raw: str = metadata.get("usage_tips") or "{}"
model_description: str = metadata.get("modelDescription") or ""
base_model: str = metadata.get("base_model") or ""
preview_url: str = metadata.get("preview_url") or ""
confidence: str = metadata.get("_llm_confidence") or ""
# --- base_model ---
base_model_valid = base_model in SUPPORTED_BASE_MODELS
base_model_filled = bool(base_model) and base_model != "Unknown"
# --- trigger_words (trainedWords) ---
triggers_valid = _is_valid_trigger_words(trained_words)
# --- short_description (civitai.description) ---
desc_filled = len(short_desc.strip()) >= _MIN_DESC_LENGTH
# --- tags ---
tags_valid = _is_valid_tags(tags)
tags_priority_coverage = _tag_priority_coverage(tags)
tags_no_technical = (
sum(1 for t in tags if t.strip().lower() not in _TECH_TAGS) >= _MIN_TAGS
if tags else False
)
# --- notes ---
notes_filled = len(notes.strip()) >= _MIN_NOTES_LENGTH
# --- usage_tips ---
usage_tips_valid = False
if usage_tips_raw.strip() and usage_tips_raw.strip() != "{}":
try:
parsed = json.loads(usage_tips_raw)
if isinstance(parsed, dict) and len(parsed) > 0:
usage_tips_valid = True
except (json.JSONDecodeError, TypeError):
pass
# --- modelDescription (README → HTML) ---
desc_html_filled = len(model_description.strip()) > 100
# --- preview_url ---
preview_filled = bool(preview_url) and os.path.exists(preview_url)
# ------------------------------------------------------------------
# Composite score (0-100)
# ------------------------------------------------------------------
field_scores = {
"base_model": _score_bool(base_model_filled and base_model_valid, weight=15),
"trigger_words": _score_bool(triggers_valid, weight=15),
"short_description": _score_bool(desc_filled, weight=10),
"tags": _score_bool(tags_valid, weight=15),
"tags_priority_coverage": _score_continuous(tags_priority_coverage, weight=5),
"notes": _score_bool(notes_filled, weight=5),
"usage_tips": _score_bool(usage_tips_valid, weight=5),
"modelDescription_html": _score_bool(desc_html_filled, weight=10),
"preview_downloaded": _score_bool(preview_filled, weight=10),
}
# Deduct points for enrichment-level failures
penalty = 0
if enrichment_errors:
penalty += 10
if not enrichment_success:
penalty += 20
total_raw = sum(field_scores.values())
total = max(0, min(100, total_raw - penalty))
# ------------------------------------------------------------------
# Flagged issues
# ------------------------------------------------------------------
issues: List[str] = []
if not base_model_filled:
issues.append("base_model is empty or 'Unknown'")
elif not base_model_valid:
issues.append(f"base_model '{base_model}' not in SUPPORTED_BASE_MODELS")
if not triggers_valid:
issues.append("trigger_words are missing or contain only placeholders")
if not desc_filled:
issues.append("short_description is too short or empty")
if not tags_valid:
issues.append("tags are missing, too few, or purely technical")
if tags_valid and tags_priority_coverage < 0.5:
issues.append("tags have low overlap with priority_tags (< 50%)")
if not notes_filled:
issues.append("notes are too short or empty")
if not usage_tips_valid:
issues.append("usage_tips is empty or invalid JSON")
if not desc_html_filled:
issues.append("modelDescription is too short (README may not have been converted)")
if not preview_filled:
issues.append("preview image not downloaded (URL missing or download failed)")
return {
"repo_id": repo_id,
"model_path": model_path,
"enrichment_success": enrichment_success,
"total_score": total,
"field_scores": field_scores,
"issues": issues,
"confidence_from_llm": confidence,
"raw_values": {
"base_model": base_model,
"trigger_words": trained_words,
"short_description": short_desc,
"tags": tags,
"notes": notes,
"usage_tips": usage_tips_raw,
"preview_url": preview_url,
"has_modelDescription": len(model_description) > 0,
},
}
def _score_bool(condition: bool, weight: int = 10) -> int:
return weight if condition else 0
def _score_continuous(value: float, weight: int = 10) -> int:
"""Linear interpolation: value 0.0 → 0, value 1.0 → *weight*."""
return int(round(value * weight))
# ---------------------------------------------------------------------------
# Batch evaluation
# ---------------------------------------------------------------------------
def evaluate_batch(
enriched: List[Dict[str, Any]],
) -> List[ScoreRecord]:
"""Evaluate a list of enrichment results.
Each entry in *enriched* should have keys:
``repo_id``, ``model_path``, ``metadata`` (the enriched dict),
``success``, ``errors``.
"""
scores: List[ScoreRecord] = []
for entry in enriched:
record = evaluate_model(
metadata=entry.get("metadata", {}),
model_path=entry.get("model_path", ""),
repo_id=entry.get("repo_id", ""),
enrichment_success=entry.get("success", False),
enrichment_errors=entry.get("errors", []),
)
scores.append(record)
return scores
# ---------------------------------------------------------------------------
# Aggregate statistics
# ---------------------------------------------------------------------------
def aggregate_scores(scores: List[ScoreRecord]) -> Dict[str, Any]:
"""Compute aggregate stats across all scored models."""
n = len(scores)
if n == 0:
return {"error": "no scores to aggregate"}
field_names = [
"base_model", "trigger_words", "short_description", "tags",
"tags_priority_coverage", "notes", "usage_tips",
"modelDescription_html", "preview_downloaded",
]
possible = {f: 15 if f == "base_model" or f == "trigger_words" or f == "tags" else
10 if f == "short_description" or f == "modelDescription_html" or f == "preview_downloaded" else
5
for f in field_names}
# Per-field aggregate
field_agg: Dict[str, Any] = {}
for fn in field_names:
vals = [s["field_scores"].get(fn, 0) for s in scores]
max_per_field = possible[fn]
field_agg[fn] = {
"mean": round(sum(vals) / n, 1) if n else 0,
"fill_rate_pct": round(
sum(1 for v in vals if v >= max_per_field) / n * 100, 1
) if n else 0.0,
"partial_rate_pct": round(
sum(1 for v in vals if 0 < v < max_per_field) / n * 100, 1
) if n else 0.0,
"empty_rate_pct": round(
sum(1 for v in vals if v == 0) / n * 100, 1
) if n else 0.0,
}
# Total score distribution
total_scores = [s["total_score"] for s in scores]
total_agg = {
"mean": round(sum(total_scores) / n, 1) if n else 0,
"median": _median(total_scores),
"min": min(total_scores) if total_scores else 0,
"max": max(total_scores) if total_scores else 0,
"bins": {
"excellent_80+": sum(1 for s in total_scores if s >= 80),
"good_60_79": sum(1 for s in total_scores if 60 <= s < 80),
"fair_40_59": sum(1 for s in total_scores if 40 <= s < 60),
"poor_20_39": sum(1 for s in total_scores if 20 <= s < 40),
"bad_0_19": sum(1 for s in total_scores if s < 20),
},
}
# Issue frequency
issue_counter: Dict[str, int] = {}
for s in scores:
for issue in s["issues"]:
issue_counter[issue] = issue_counter.get(issue, 0) + 1
top_issues = sorted(issue_counter.items(), key=lambda x: -x[1])
# Confidence distribution
conf_counter: Dict[str, int] = {"high": 0, "medium": 0, "low": 0, "": 0}
for s in scores:
c = (s.get("confidence_from_llm") or "").strip().lower()
if c in conf_counter:
conf_counter[c] += 1
else:
conf_counter[""] += 1
# Success / timeout / failure stats
success_count = sum(1 for s in scores if s["enrichment_success"])
fail_count = n - success_count
return {
"model_count": n,
"success_count": success_count,
"fail_count": fail_count,
"total_score": total_agg,
"field_aggregates": field_agg,
"top_issues": top_issues[:15],
"confidence_distribution": conf_counter,
}
def _median(values: List[float]) -> float:
if not values:
return 0.0
sorted_v = sorted(values)
m = len(sorted_v) // 2
if len(sorted_v) % 2 == 0:
return round((sorted_v[m - 1] + sorted_v[m]) / 2, 1)
return round(sorted_v[m], 1)
@@ -0,0 +1,202 @@
"""Construct initial ``.metadata.json`` sidecars for HF model repos.
Each HF repo + safetensors pair gets a minimal metadata file no real model
file is needed. The enrichment pipeline reads only the sidecar.
Data format (one line per entry)::
repo_id, model_name.safetensors
"""
from __future__ import annotations
import json
import logging
import os
from typing import Any, Dict, List, Tuple
from .config import CIVITAI_MODEL_TAGS
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Data types
# ---------------------------------------------------------------------------
# A validated entry parsed from the models file:
# (repo_id, safetensors_name)
RepoEntry = Tuple[str, str]
def load_repo_ids(path: str, max_models: int | None = None) -> List[RepoEntry]:
"""Read ``repo_id, safetensors_name`` pairs from *path*.
Format (one per line, blanks and ``#`` comments ignored)::
user/repo-name, lora_zimage_turbo_myjs_alpha01.safetensors
Returns a list of ``(repo_id, safetensors_name)`` tuples.
"""
path = os.path.expanduser(path)
if not os.path.exists(path):
raise FileNotFoundError(f"Models file not found: {path}")
entries: List[RepoEntry] = []
with open(path, "r", encoding="utf-8") as fh:
for raw_line in fh:
line = raw_line.strip()
if not line or line.startswith("#"):
continue
# Split on the first comma
if "," not in line:
logger.warning("Skipping malformed line (no comma): %s", raw_line.rstrip())
continue
repo_id, safetensors_name = [part.strip() for part in line.split(",", 1)]
if not repo_id or not safetensors_name:
logger.warning("Skipping malformed line (empty fields): %s", raw_line.rstrip())
continue
if not safetensors_name.lower().endswith(".safetensors"):
logger.warning(
"Skipping line — safetensors_name doesn't end with .safetensors: %s",
raw_line.rstrip(),
)
continue
entries.append((repo_id, safetensors_name))
if max_models is not None and max_models > 0:
entries = entries[:max_models]
logger.info("Loaded %d HF repo entries from %s", len(entries), path)
return entries
def sanitize_repo_id(repo_id: str) -> str:
"""Turn ``user/repo-name`` into a safe directory name."""
return repo_id.replace("/", "__").replace(".", "_")
def build_model_dir(output_dir: str, repo_id: str) -> str:
"""Return the per-model working directory."""
return os.path.join(output_dir, "models", sanitize_repo_id(repo_id))
def build_model_path(model_dir: str, safetensors_name: str) -> str:
"""Return the model file path using the real safetensors filename."""
return os.path.join(model_dir, safetensors_name)
def build_metadata_path(model_path: str) -> str:
"""Return the sidecar path for a model file.
This MUST match the convention used by ``MetadataManager`` /
``apply_metadata_updates``, which derives the sidecar path via
``os.path.splitext(model_path)[0] + '.metadata.json'``.
For a model file ``lora_x.safetensors`` the sidecar is
``lora_x.metadata.json`` *not* ``lora_x.safetensors.metadata.json``.
"""
return f"{os.path.splitext(model_path)[0]}.metadata.json"
def create_initial_metadata(
output_dir: str,
repo_id: str,
safetensors_name: str,
) -> str:
"""Write a minimal ``.metadata.json`` for *repo_id* + *safetensors_name*.
Args:
output_dir: Root output directory.
repo_id: HuggingFace repo identifier (``user/repo``).
safetensors_name: The specific model file name (e.g.
``lora_zimage_turbo_myjs_alpha01.safetensors``).
Returns the **model path** (the ``.safetensors`` path whose sidecar was
written). The caller passes this path to ``AgentService.execute_skill``.
The basename (filename without extension) will match the real model file,
so ``extract_relevant_section`` can reliably match against the README.
"""
model_dir = build_model_dir(output_dir, repo_id)
os.makedirs(model_dir, exist_ok=True)
model_path = build_model_path(model_dir, safetensors_name)
metadata_path = build_metadata_path(model_path)
hf_url = f"https://huggingface.co/{repo_id}"
file_name = safetensors_name
metadata: Dict[str, Any] = {
"file_name": file_name,
"model_name": safetensors_name,
"file_path": model_path.replace(os.sep, "/"),
"size": 0,
"modified": 0,
"sha256": "",
"base_model": "Unknown",
"preview_url": "",
"preview_nsfw_level": 0,
"notes": "",
"from_civitai": False,
"civitai": {},
"tags": [],
"modelDescription": "",
"civitai_deleted": False,
"favorite": False,
"exclude": False,
"db_checked": False,
"skip_metadata_refresh": False,
"metadata_source": "",
"last_checked_at": 0,
"hash_status": "completed",
"trainedWords": [],
"hf_url": hf_url,
"usage_tips": "{}",
}
with open(metadata_path, "w", encoding="utf-8") as fh:
json.dump(metadata, fh, indent=2, ensure_ascii=False)
logger.debug("Created initial metadata for %s -> %s", repo_id, metadata_path)
return model_path
def create_all_initial_metadata(
entries: List[RepoEntry],
output_dir: str,
*,
skip_existing: bool = True,
) -> Tuple[List[str], List[str]]:
"""Create initial metadata for every repo entry.
Args:
entries: List of ``(repo_id, safetensors_name)`` tuples.
output_dir: Root output directory.
skip_existing: If True, skip repos whose metadata already exists.
Returns:
A tuple ``(model_paths, repo_ids)`` two parallel lists in the same
order as *entries*. This keeps downstream code (enrichment runner,
evaluation engine) unchanged.
"""
model_paths: List[str] = []
repo_ids: List[str] = []
for repo_id, safetensors_name in entries:
model_dir = build_model_dir(output_dir, repo_id)
model_path = build_model_path(model_dir, safetensors_name)
metadata_path = build_metadata_path(model_path)
if skip_existing and os.path.exists(metadata_path):
model_paths.append(model_path)
repo_ids.append(repo_id)
continue
model_paths.append(create_initial_metadata(output_dir, repo_id, safetensors_name))
repo_ids.append(repo_id)
logger.info(
"Constructed initial metadata for %d/%d repos",
len(model_paths),
len(entries),
)
return model_paths, repo_ids
@@ -0,0 +1,467 @@
"""Preprocessing audit for the HF metadata enrichment validation pipeline.
Phase 1.5 runs between Phase 1 (metadata creation) and Phase 2 (enrichment).
Audits the README preprocessing pipeline (section extraction + cleaning)
for each repo in the dataset, capturing intermediate outputs so we can
distinguish between:
(A) Preprocessing failed LLM never saw the right content
(B) Preprocessing succeeded LLM/prompt needs improvement
This prevents wasted effort optimizing prompts when the actual problem is
that ``extract_relevant_section`` or ``clean_readme_for_llm`` removed or
misaligned the content the LLM needed.
"""
from __future__ import annotations
import asyncio
import json
import logging
import os
import re
import time
from dataclasses import dataclass, field, asdict
from typing import Any, Dict, List, Tuple
import aiohttp
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Audit record
# ---------------------------------------------------------------------------
@dataclass
class AuditRecord:
"""Preprocessing audit for a single repo entry."""
# Identity
repo_id: str
safetensors_name: str
basename: str # filename without .safetensors
# Raw README stats
raw_readme_length: int
raw_readme_line_count: int
has_yaml_frontmatter: bool
yaml_has_base_model: bool
yaml_has_tags: bool
# Section extraction
section_extraction_activated: bool # output < 95% of input length
section_length: int
section_line_count: int
basename_in_section: bool # basename appears in extracted section text
# Cleaning
cleaned_length: int
cleaned_line_count: int
compression_pct: float # (1 - cleaned/raw) * 100
# Widget section (stripped by _strip_widget_section)
widget_section_found: bool
widget_section_length: int
# Flags (list of anomaly descriptions)
flags: List[str] = field(default_factory=list)
# Local file path to the saved raw README (for cross-reference)
readme_file: str = ""
# Staged intermediate output for report detail
raw_readme_preview: str = "" # first 200 chars
section_preview: str = "" # first 300 chars
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
_HF_RAW_URL = "https://huggingface.co/{repo_id}/raw/main/README.md"
# Thresholds for flagging
_SECTION_ACTIVATION_RATIO = 0.95
_MIN_CLEANED_LENGTH = 100
_MAX_COMPRESSION_PCT = 99.0
_MIN_SECTION_LINES = 3
# ---------------------------------------------------------------------------
# Module loader — bypasses parent-package __init__ that imports ComfyUI
# ---------------------------------------------------------------------------
_readme_processor_module = None
def _load_readme_processor():
"""Import ``readme_processor`` without triggering ``folder_paths`` import.
The normal import path (``py.services.agent.skills.enrich_hf_metadata.
readme_processor``) triggers ``py.services.agent.__init__`` which
imports ``agent_service.py`` ``py/config.py`` ComfyUI's
``folder_paths``, which is not available in standalone mode.
"""
global _readme_processor_module
if _readme_processor_module is not None:
return _readme_processor_module
import importlib.util
_RP_PATH = os.path.join(
os.path.dirname(__file__), # tests/enrich_hf_validation/
"..", "..",
"py", "services", "agent", "skills", "enrich_hf_metadata",
"readme_processor.py",
)
rp_path = os.path.normpath(_RP_PATH)
if not os.path.exists(rp_path):
logger.error("readme_processor.py not found at %s", rp_path)
return None
spec = importlib.util.spec_from_file_location(
"readme_processor", rp_path,
)
if spec is None or spec.loader is None:
logger.error("Could not create spec for readme_processor.py")
return None
mod = importlib.util.module_from_spec(spec)
try:
spec.loader.exec_module(mod)
except Exception as exc:
logger.error("Failed to load readme_processor.py: %s", exc)
return None
_readme_processor_module = mod
return mod
# ---------------------------------------------------------------------------
# HF README fetcher
# ---------------------------------------------------------------------------
async def _fetch_readme(repo_id: str, session: aiohttp.ClientSession) -> str:
"""Fetch the raw README.md from HuggingFace."""
url = _HF_RAW_URL.format(repo_id=repo_id)
try:
async with session.get(url, timeout=aiohttp.ClientTimeout(total=30)) as resp:
if resp.status == 200:
return await resp.text()
logger.warning("Failed to fetch README for %s: HTTP %d", repo_id, resp.status)
return ""
except (asyncio.TimeoutError, aiohttp.ClientError) as exc:
logger.warning("Failed to fetch README for %s: %s", repo_id, exc)
return ""
# ---------------------------------------------------------------------------
# Analysis helpers
# ---------------------------------------------------------------------------
def _has_yaml_frontmatter(text: str) -> bool:
return bool(text.strip().startswith("---"))
def _extract_yaml_field(text: str, field: str) -> bool:
"""Check if the given YAML field exists in the frontmatter."""
lines = text.split("\n")
if not lines or not lines[0].strip().startswith("---"):
return False
end = 1
while end < len(lines):
if lines[end].strip().startswith("---"):
break
end += 1
if end >= len(lines):
return False
frontmatter = "\n".join(lines[1:end])
pattern = rf"^{field}:"
return bool(re.search(pattern, frontmatter, re.MULTILINE))
def _find_widget_section_length(text: str) -> int:
"""Find the ``widget:`` YAML section and return its length (0 if none)."""
if not _has_yaml_frontmatter(text):
return 0
frontmatter_end = text.find("---", 3)
if frontmatter_end == -1:
return 0
frontmatter = text[3:frontmatter_end]
# Match widget: through to the next top-level key or frontmatter end
m = re.search(r"\nwidget:", frontmatter)
if not m:
return 0
# Length from widget: to end of frontmatter (the next \n\w+: or \n---)
return len(frontmatter[m.start():])
# ---------------------------------------------------------------------------
# Core auditor
# ---------------------------------------------------------------------------
async def run_audit(
entries: List[Tuple[str, str]],
*,
concurrency: int = 10,
readmes_dir: str | None = None,
) -> Tuple[List[AuditRecord], Dict[str, Any]]:
"""Run the preprocessing audit over all repo entries.
Args:
entries: List of ``(repo_id, safetensors_name)``.
concurrency: Max parallel fetches to HuggingFace.
readmes_dir: If set, saves each fetched README as
``{sanitized_repo_id}.md`` in this directory for offline
cross-reference against audit results.
Returns:
Tuple of ``(records, summary)`` where *summary* is a dict with
aggregate statistics.
"""
semaphore = asyncio.Semaphore(concurrency)
records: List[AuditRecord] = []
flag_counter: Dict[str, int] = {}
if readmes_dir:
os.makedirs(readmes_dir, exist_ok=True)
connector = aiohttp.TCPConnector(limit=concurrency)
async with aiohttp.ClientSession(connector=connector) as session:
tasks = [_audit_one(entry, session, semaphore, readmes_dir=readmes_dir) for entry in entries]
gathered = await asyncio.gather(*tasks, return_exceptions=True)
for entry, result in zip(entries, gathered):
if isinstance(result, Exception):
logger.error("Audit failed for %s: %s", entry[0], result)
records.append(
AuditRecord(
repo_id=entry[0],
safetensors_name=entry[1],
basename=os.path.splitext(entry[1])[0],
raw_readme_length=0,
raw_readme_line_count=0,
has_yaml_frontmatter=False,
yaml_has_base_model=False,
yaml_has_tags=False,
section_extraction_activated=False,
section_length=0,
section_line_count=0,
basename_in_section=False,
cleaned_length=0,
cleaned_line_count=0,
compression_pct=0.0,
widget_section_found=False,
widget_section_length=0,
readme_file="",
flags=[f"Audit exception: {result}"],
)
)
continue
# The continue above ensures result is AuditRecord here
assert isinstance(result, AuditRecord)
records.append(result)
for flag in result.flags:
flag_counter[flag] = flag_counter.get(flag, 0) + 1
summary = _build_summary(records, flag_counter)
return records, summary
def _sanitize_repo_id(repo_id: str) -> str:
"""Turn ``user/repo-name`` into a safe filename."""
return repo_id.replace("/", "__").replace(".", "_")
async def _audit_one(
entry: Tuple[str, str],
session: aiohttp.ClientSession,
semaphore: asyncio.Semaphore,
*,
readmes_dir: str | None = None,
) -> AuditRecord:
"""Audit a single repo entry."""
repo_id, safetensors_name = entry
basename = os.path.splitext(safetensors_name)[0]
async with semaphore:
# Import production preprocessing functions.
# Use importlib to bypass py.services.agent.__init__ which triggers
# ComfyUI's folder_paths module (not available in standalone mode).
_rp = _load_readme_processor()
if _rp is None:
return AuditRecord(
repo_id=repo_id,
safetensors_name=safetensors_name,
basename=basename,
raw_readme_length=0, raw_readme_line_count=0,
has_yaml_frontmatter=False, yaml_has_base_model=False, yaml_has_tags=False,
readme_file="",
section_extraction_activated=False, section_length=0, section_line_count=0,
basename_in_section=False, cleaned_length=0, cleaned_line_count=0,
compression_pct=0.0, widget_section_found=False, widget_section_length=0,
flags=["IMPORT_FAILED"],
)
clean_readme_for_llm = _rp.clean_readme_for_llm
extract_relevant_section = _rp.extract_relevant_section
# Step 1: Fetch the raw README
raw_text = await _fetch_readme(repo_id, session)
if not raw_text:
return AuditRecord(
repo_id=repo_id,
safetensors_name=safetensors_name,
basename=basename,
raw_readme_length=0,
raw_readme_line_count=0,
has_yaml_frontmatter=False,
yaml_has_base_model=False,
yaml_has_tags=False,
section_extraction_activated=False,
section_length=0,
section_line_count=0,
basename_in_section=False,
readme_file="",
cleaned_length=0,
cleaned_line_count=0,
compression_pct=0.0,
widget_section_found=False,
widget_section_length=0,
flags=["README_FETCH_FAILED"],
)
# Save the raw README to disk for offline cross-reference
readme_path = ""
if readmes_dir:
safe_name = _sanitize_repo_id(repo_id)
readme_path = os.path.join(readmes_dir, f"{safe_name}.md")
try:
with open(readme_path, "w", encoding="utf-8") as fh:
fh.write(raw_text)
except OSError as exc:
logger.warning("Failed to save README for %s: %s", repo_id, exc)
readme_path = ""
raw_lines = raw_text.split("\n")
raw_len = len(raw_text)
raw_line_count = len(raw_lines)
# Step 2: Analyze raw README
yaml_fm = _has_yaml_frontmatter(raw_text)
yaml_has_bm = _extract_yaml_field(raw_text, "base_model") if yaml_fm else False
yaml_has_tg = _extract_yaml_field(raw_text, "tags") if yaml_fm else False
widget_len = _find_widget_section_length(raw_text)
# Step 3: Section extraction
section = extract_relevant_section(raw_text, basename)
section_len = len(section)
section_line_count = len(section.split("\n"))
section_activated = section_len < raw_len * _SECTION_ACTIVATION_RATIO
basename_in_sec = basename.lower() in section.lower()
# Step 4: Cleaning for LLM
cleaned = clean_readme_for_llm(section)
cleaned_len = len(cleaned)
cleaned_line_count = len(cleaned.split("\n"))
compression_pct = round((1 - cleaned_len / raw_len) * 100, 1) if raw_len else 0.0
# Step 5: Flag anomalies
flags: List[str] = []
if not raw_text.strip():
flags.append("README_EMPTY")
if not yaml_fm:
flags.append("NO_YAML_FRONTMATTER")
if not section_activated:
# Check if basename is extremely short/generic (likely synthetic)
if len(basename) <= 5:
flags.append("BASENAME_TOO_SHORT_SECTION_NOT_EXPECTED")
else:
flags.append("SECTION_EXTRACTION_NOT_ACTIVATED")
elif not basename_in_sec:
flags.append("BASENAME_NOT_IN_EXTRACTED_SECTION")
if widget_len == 0:
# Not necessarily a problem — many repos lack a widget section
pass
if cleaned_len < _MIN_CLEANED_LENGTH:
flags.append("CLEANED_README_TOO_SHORT")
if compression_pct > _MAX_COMPRESSION_PCT:
flags.append("EXTREME_COMPRESSION")
if section_activated and section_line_count < _MIN_SECTION_LINES:
flags.append("SECTION_TOO_SMALL")
return AuditRecord(
repo_id=repo_id,
safetensors_name=safetensors_name,
basename=basename,
raw_readme_length=raw_len,
raw_readme_line_count=raw_line_count,
has_yaml_frontmatter=yaml_fm,
yaml_has_base_model=yaml_has_bm,
yaml_has_tags=yaml_has_tg,
section_extraction_activated=section_activated,
section_length=section_len,
section_line_count=section_line_count,
basename_in_section=basename_in_sec,
cleaned_length=cleaned_len,
cleaned_line_count=cleaned_line_count,
compression_pct=compression_pct,
widget_section_found=widget_len > 0,
widget_section_length=widget_len,
readme_file=readme_path,
flags=flags,
raw_readme_preview=raw_text[:200],
section_preview=section[:300],
)
def _build_summary(
records: List[AuditRecord],
flag_counter: Dict[str, int],
) -> Dict[str, Any]:
"""Aggregate audit statistics."""
n = len(records)
if n == 0:
return {"error": "no records", "model_count": 0}
activated = sum(1 for r in records if r.section_extraction_activated)
basename_hit = sum(1 for r in records if r.basename_in_section)
with_yaml = sum(1 for r in records if r.has_yaml_frontmatter)
with_widget = sum(1 for r in records if r.widget_section_found)
fetch_failed = sum(1 for r in records if "README_FETCH_FAILED" in r.flags)
avg_compression = round(
sum(r.compression_pct for r in records if r.raw_readme_length > 0) / max(n - fetch_failed, 1),
1,
)
avg_cleaned = round(
sum(r.cleaned_length for r in records if r.raw_readme_length > 0) / max(n - fetch_failed, 1),
)
top_flags = sorted(flag_counter.items(), key=lambda x: -x[1])[:10]
return {
"model_count": n,
"fetch_failed_count": fetch_failed,
"section_extraction_activated": activated,
"section_extraction_pct": round(activated / max(n - fetch_failed, 1) * 100, 1),
"basename_in_section": basename_hit,
"basename_in_section_pct": round(basename_hit / max(n - fetch_failed, 1) * 100, 1),
"with_yaml_frontmatter": with_yaml,
"with_yaml_frontmatter_pct": round(with_yaml / max(n - fetch_failed, 1) * 100, 1),
"with_widget_section": with_widget,
"avg_compression_pct": avg_compression,
"avg_cleaned_length": avg_cleaned,
"top_flags": top_flags,
}
def audit_records_to_serializable(records: List[AuditRecord]) -> List[Dict[str, Any]]:
"""Convert AuditRecord dataclasses to plain dicts for JSON serialization."""
return [asdict(r) for r in records]
@@ -0,0 +1,391 @@
"""Generate structured reports from evaluation results.
Produces:
1. A JSON data dump (``report.json``) with all scores and aggregations.
2. A human-readable Markdown report (``report.md``) with summary stats,
issue patterns, and actionable optimisation suggestions.
"""
from __future__ import annotations
import json
import logging
import os
from datetime import datetime
from typing import Any, Dict, List
from .config import SUPPORTED_BASE_MODELS
from .evaluation_engine import ScoreRecord
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Markdown report
# ---------------------------------------------------------------------------
def _fmt_pct(value: float) -> str:
return f"{value:.1f}%"
def _bar(value: float, width: int = 20) -> str:
filled = int(round(value / 100 * width))
return "" * filled + "" * (width - filled)
def generate_optimisation_suggestions(
agg: Dict[str, Any],
scores: List[ScoreRecord],
) -> List[str]:
"""Analyse evaluation results and produce concrete suggestions."""
suggestions: List[str] = []
fa = agg.get("field_aggregates", {})
# --- base_model ---
bm = fa.get("base_model", {})
if bm and bm.get("empty_rate_pct", 0) > 30:
suggestions.append(
"- **base_model 空置率高 ({:.0f}%)**: 多数 HF 模型卡片未在 YAML frontmatter 中声明 "
"`base_model:` 字段,LLM 无法推断。可考虑在 prompt 中增加 \"look at the model file name "
"for clues\" 的引导,或在后处理中增加基于文件名规则的 fallback 猜测。".format(
bm.get("empty_rate_pct", 0)
)
)
bm_invalid = sum(
1
for s in scores
if s["raw_values"]["base_model"]
and s["raw_values"]["base_model"] != "Unknown"
and s["raw_values"]["base_model"] not in set(SUPPORTED_BASE_MODELS)
)
if bm_invalid > 5:
suggestions.append(
"- **base_model 含非标准值 ({} 个)**: LLM 输出了未在当前生产系统的 base model 列表 "
"中的名称。建议在 prompt 中强调 \"Use EXACTLY one name from the list\" 并在 "
"`PostProcessor` 中加一层验证过滤,非标准值直接丢弃。".format(bm_invalid)
)
# --- trigger_words ---
tw = fa.get("trigger_words", {})
if tw and tw.get("empty_rate_pct", 0) > 40:
suggestions.append(
"- **trigger_words 空置率高 ({:.0f}%)**: 大量 HF 模型卡没有明确的 "
"`instance_prompt:` 或 trigger word 说明。当前 prompt 已覆盖常见模式。若确认这些模型确实"
"没有 trigger words(例如 style lora),空数组是正确结果,不需优化。".format(
tw.get("empty_rate_pct", 0)
)
)
# --- tags ---
tag = fa.get("tags", {})
if tag and tag.get("empty_rate_pct", 0) > 30:
suggestions.append(
"- **tags 空置率高 ({:.0f}%)**: 当前 prompt 要求 tags 必须与 "
"`priority_tags`CIVITAI_MODEL_TAGS)对齐。HF 模型的标签体系与 Civitai 不同,"
"很多 model card 使用细粒度标签(如 `pokemon`、`watercolor`)而不在 priority list 中。"
"建议: 扩大 priority_tags 范围,或允许 LLM 自由生成 tags 后只做去重不做严格过滤。".format(
tag.get("empty_rate_pct", 0)
)
)
# --- tags priority coverage ---
low_coverage = sum(
1
for s in scores
if s["field_scores"].get("tags_priority_coverage", 5) < 3 # < 60% of max
and s["field_scores"].get("tags", 0) > 0
)
if low_coverage > 10:
suggestions.append(
"- **{} 个模型的 tags 与 priority_tags 匹配度低于 60%**: "
"LLM 生成了有意义但不属于 CIVITAI_MODEL_TAGS 的标签。这说明 priority_tags "
"的覆盖范围对 HF 模型不足,建议按 HF 模型的实际分布补充新类别。".format(low_coverage)
)
# --- preview ---
prev = fa.get("preview_downloaded", {})
if prev and prev.get("empty_rate_pct", 0) > 50:
suggestions.append(
"- **预览图下载成功率低 ({:.0f}%)**: 很多 HF 模型卡没有 embed 图片(仅使用 YAML widget "
"或 external link)。当前 `readme_processor.py` 的 `extract_gallery_images` 和 "
"`extract_gallery_table_images` 已覆盖了多数场景。若预览图不重要,可降低此字段权重。".format(
prev.get("empty_rate_pct", 0)
)
)
# --- usage_tips ---
ut = fa.get("usage_tips", {})
if ut and ut.get("empty_rate_pct", 0) > 70:
suggestions.append(
"- **usage_tips 空置率极高 ({:.0f}%)**: 这是预期行为。HF 模型卡通常不包含 LoRA "
"强度/CLIP skip 等结构化参数。当前提取策略已合理。若需要可用数据,"
"可以考虑使用模型类型的通用默认值。".format(
ut.get("empty_rate_pct", 0)
)
)
# --- short_description ---
sd = fa.get("short_description", {})
if sd and sd.get("empty_rate_pct", 0) > 40:
suggestions.append(
"- **short_description 空置率 ({:.0f}%)**: 部分 HF 模型卡 README 内容极少(仅含标签和训练参数)。".format(
sd.get("empty_rate_pct", 0)
)
)
if not suggestions:
suggestions.append("- 未发现明显问题模式,各字段填充率均在可接受范围。")
return suggestions
def generate_markdown_report(
agg: Dict[str, Any],
scores: List[ScoreRecord],
output_dir: str,
duration_summary: Dict[str, Any] | None = None,
*,
audit_summary: Dict[str, Any] | None = None,
config_warnings: List[str] | None = None,
) -> str:
"""Write ``report.md`` and return its content.
Args:
agg: Aggregate evaluation scores.
scores: Per-model evaluation records.
output_dir: Output directory for the report file.
duration_summary: Optional timing statistics.
audit_summary: Optional preprocessing audit summary (Phase 1.5).
config_warnings: Optional LLM config consistency warnings.
"""
lines: List[str] = []
def wl(text: str = "") -> None:
lines.append(text)
wl("# HF Metadata Enrichment Validation Report")
wl()
wl(f"Generated: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")
wl(f"Models evaluated: **{agg.get('model_count', 0)}**")
wl(f"Successful enrichments: **{agg.get('success_count', 0)}**")
wl(f"Failures: **{agg.get('fail_count', 0)}**")
wl()
# ---- Preprocessing Audit Section ----
if audit_summary and audit_summary.get("model_count", 0) > 0:
wl("## Preprocessing Audit")
wl()
wl(f"| Metric | Value |")
wl(f"|--------|-------|")
wl(f"| Models audited | {audit_summary.get('model_count', 0)} |")
wl(f"| README fetch failed | {audit_summary.get('fetch_failed_count', 0)} |")
wl(f"| Section extraction activated | {_fmt_pct(audit_summary.get('section_extraction_pct', 0))} |")
wl(f"| Basename found in section | {_fmt_pct(audit_summary.get('basename_in_section_pct', 0))} |")
wl(f"| Has YAML frontmatter | {_fmt_pct(audit_summary.get('with_yaml_frontmatter_pct', 0))} |")
wl(f"| Has YAML widget section | {_fmt_pct(audit_summary.get('with_widget_section', 0))} |")
wl(f"| Avg README compression | {audit_summary.get('avg_compression_pct', 0)}% |")
wl(f"| Avg cleaned length | {audit_summary.get('avg_cleaned_length', 0)} chars |")
wl()
if audit_summary.get("top_flags"):
wl("### Audit Flags (most frequent)")
wl()
for flag, count in audit_summary["top_flags"]:
wl(f"- **{flag}**: {count}x")
wl()
wl("**Interpretation:**")
wl()
act_pct = audit_summary.get("section_extraction_pct", 0)
if act_pct < 50:
wl(
"- ⚠️ Section extraction activated for fewer than 50% of repos. "
"This may indicate the basename doesn't match README content, or the "
"repos are mostly single-model (where full README is expected)."
)
else:
wl(
"- ✅ Section extraction is working for most repos — the LLM is "
"receiving focused README sections."
)
if audit_summary.get("basename_in_section_pct", 100) < 80:
wl(
"- ⚠️ The safetensors basename was NOT found in the extracted section "
"for many repos. This could mean the section extraction matched the wrong "
"section, or the README doesn't explicitly reference the filename."
)
wl()
# ---- Config warnings ----
if config_warnings:
wl("## ⚠️ Configuration Warnings")
wl()
for w in config_warnings:
wl(f"- {w}")
wl()
# ---- Duration ----
if duration_summary:
wl("## Timing")
wl()
wl(f"- Total wall time: **{duration_summary.get('total_wall_s', 0):.0f} s** ")
wl(f" ({duration_summary.get('total_wall_s', 0) / 60:.1f} min)")
wl(f"- Mean per model: **{duration_summary.get('mean_s', 0):.1f} s**")
wl(f"- Median per model: **{duration_summary.get('median_s', 0):.1f} s**")
wl(f"- Fastest: **{duration_summary.get('min_s', 0):.1f} s**")
wl(f"- Slowest: **{duration_summary.get('max_s', 0):.1f} s**")
wl()
# ---- Overall score ----
ts = agg.get("total_score", {})
wl("## Overall Score Distribution (0100)")
wl()
wl(f"| Metric | Value |")
wl(f"|--------|-------|")
wl(f"| Mean | {ts.get('mean', 'N/A')} |")
wl(f"| Median | {ts.get('median', 'N/A')} |")
wl(f"| Min | {ts.get('min', 'N/A')} |")
wl(f"| Max | {ts.get('max', 'N/A')} |")
wl()
for label, key in [
("Excellent (≥80)", "excellent_80+"),
("Good (6079)", "good_60_79"),
("Fair (4059)", "fair_40_59"),
("Poor (2039)", "poor_20_39"),
("Bad (<20)", "bad_0_19"),
]:
count = ts.get("bins", {}).get(key, 0)
pct = count / agg["model_count"] * 100 if agg["model_count"] else 0
wl(f"- **{label}**: {count} models ({_fmt_pct(pct)})")
wl()
# ---- Per-field aggregates ----
wl("## Per-Field Completeness")
wl()
wl("| Field | Mean Score | Fill Rate | Empty Rate |")
wl("|-------|-----------:|----------:|-----------:|")
fa = agg.get("field_aggregates", {})
for fn in [
"base_model", "trigger_words", "short_description", "tags",
"tags_priority_coverage", "notes", "usage_tips",
"modelDescription_html", "preview_downloaded",
]:
f = fa.get(fn, {})
if not f:
continue
wl(
f"| {fn} "
f"| {f.get('mean', 'N/A')} "
f"| {_fmt_pct(f.get('fill_rate_pct', 0))} "
f"| {_fmt_pct(f.get('empty_rate_pct', 0))} |"
)
wl()
# ---- Confidence distribution ----
wl("## LLM Confidence Distribution")
wl()
cd = agg.get("confidence_distribution", {})
total_conf = sum(cd.values()) or 1
for level in ["high", "medium", "low", ""]:
count = cd.get(level, 0)
label = level if level else "(not reported)"
pct = count / total_conf * 100
bar = _bar(pct)
wl(f"- **{label}**: {count} {bar} {_fmt_pct(pct)}")
wl()
# ---- Top issues ----
wl("## Most Frequent Issues")
wl()
for issue, count in agg.get("top_issues", []):
pct = count / agg["model_count"] * 100 if agg["model_count"] else 0
wl(f"- **{issue}** — {count}/{agg['model_count']} ({_fmt_pct(pct)})")
wl()
# ---- Optimisation suggestions ----
wl("## Optimisation Suggestions")
wl()
suggestions = generate_optimisation_suggestions(agg, scores)
for s in suggestions:
wl(s)
wl()
# ---- Per-model detail ----
wl("## Per-Model Detail")
wl()
wl("<details>")
wl("<summary>Click to expand</summary>")
wl()
wl("| # | Repo ID | Score | Issues | Confidence |")
wl("|---|---------|------:|--------|------------|")
for i, s in enumerate(scores, 1):
issue_count = len(s["issues"])
issue_str = (
f"{issue_count} issue(s)" if issue_count else "✓ ok"
)
wl(
f"| {i} "
f"| {s['repo_id']} "
f"| {s['total_score']} "
f"| {issue_str} "
f"| {s.get('confidence_from_llm', '') or '-'} |"
)
wl()
wl("</details>")
wl()
content = "\n".join(lines)
report_path = os.path.join(output_dir, "report.md")
with open(report_path, "w", encoding="utf-8") as fh:
fh.write(content)
logger.info("Markdown report written to %s", report_path)
return content
# ---------------------------------------------------------------------------
# JSON dump
# ---------------------------------------------------------------------------
def save_json_report(
agg: Dict[str, Any],
scores: List[ScoreRecord],
enrichment_results: List[Dict[str, Any]],
output_dir: str,
duration_summary: Dict[str, Any] | None = None,
*,
audit_summary: Dict[str, Any] | None = None,
config_warnings: List[str] | None = None,
) -> str:
"""Write ``report.json`` and return the path.
Args:
agg: Aggregate evaluation scores.
scores: Per-model evaluation records.
enrichment_results: Raw enrichment phase results.
output_dir: Output directory.
duration_summary: Optional timing statistics.
audit_summary: Optional preprocessing audit summary.
config_warnings: Optional LLM config consistency warnings.
"""
report: Dict[str, Any] = {
"metadata": {
"generated_at": datetime.now().isoformat(),
"model_count": agg.get("model_count", 0),
},
"aggregate": agg,
"timing": duration_summary or {},
"per_model_scores": scores,
"enrichment_results": enrichment_results,
}
if audit_summary:
report["preprocessing_audit"] = audit_summary
if config_warnings:
report["config_warnings"] = config_warnings
path = os.path.join(output_dir, "report.json")
with open(path, "w", encoding="utf-8") as fh:
json.dump(report, fh, indent=2, ensure_ascii=False)
logger.info("JSON report written to %s", path)
return path
@@ -0,0 +1,451 @@
#!/usr/bin/env python3
"""CLI entry point for the HF metadata enrichment validation suite.
Usage::
# Full run (44 models, serial, ~1-2 h)
python -m tests.enrich_hf_validation.run_validation \\
--output /tmp/hf_enrich_validation
# Quick smoke test with 2 models
python -m tests.enrich_hf_validation.run_validation --sample 2
# Resume from a previous partial run
python -m tests.enrich_hf_validation.run_validation --resume
# Audit preprocessing only (no LLM calls, fast)
python -m tests.enrich_hf_validation.run_validation --audit-only
# Custom settings file
python -m tests.enrich_hf_validation.run_validation \\
--settings /custom/path/settings.json
"""
from __future__ import annotations
import argparse
import asyncio
import json
import logging
import os
import sys
import time
from typing import Any, Dict, List, Tuple
# Ensure the project root is on sys.path so that ``from py import ...`` works.
_PROJECT_ROOT = os.path.normpath(
os.path.join(os.path.dirname(__file__), "..", "..")
)
if _PROJECT_ROOT not in sys.path:
sys.path.insert(0, _PROJECT_ROOT)
# Add ComfyUI root to sys.path so ``folder_paths`` can be imported.
# Project layout: ComfyUI/custom_nodes/ComfyUI-Lora-Manager/
_COMFYUI_ROOT = os.path.normpath(os.path.join(_PROJECT_ROOT, "..", ".."))
if _COMFYUI_ROOT not in sys.path:
sys.path.insert(0, _COMFYUI_ROOT)
from tests.enrich_hf_validation.config import (
init_supported_base_models,
load_settings,
)
from tests.enrich_hf_validation.metadata_constructor import (
RepoEntry,
create_all_initial_metadata,
load_repo_ids,
)
from tests.enrich_hf_validation.enrichment_runner import EnrichmentRunner
from tests.enrich_hf_validation.evaluation_engine import (
aggregate_scores,
evaluate_batch,
)
from tests.enrich_hf_validation.preprocessing_auditor import (
audit_records_to_serializable,
run_audit,
)
from tests.enrich_hf_validation.report_generator import (
generate_markdown_report,
save_json_report,
)
logger = logging.getLogger(__name__)
def _setup_logging(verbose: bool) -> None:
level = logging.DEBUG if verbose else logging.INFO
fmt = "%(asctime)s [%(levelname)s] %(name)s: %(message)s"
logging.basicConfig(level=level, format=fmt, stream=sys.stderr)
# Quiet noisy third-party loggers
for name in ("aiohttp", "asyncio", "urllib3"):
logging.getLogger(name).setLevel(logging.WARNING)
def _parse_args(argv: List[str]) -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Validate and optimise HF metadata enrichment via LLM.",
)
parser.add_argument(
"--models",
default=os.path.join(os.path.dirname(__file__), "test_data", "hf_lora_models_with_safetensors.txt"),
help="Path to the HF repo entries file (format: repo_id, model_name.safetensors per line)",
)
parser.add_argument(
"--settings",
default="~/.config/ComfyUI-LoRA-Manager/settings.json",
help="Path to LoRA Manager settings.json",
)
parser.add_argument(
"--output",
default="/tmp/hf_enrich_validation",
help="Output directory for reports and intermediate data",
)
parser.add_argument(
"--sample",
type=int,
default=0,
help="Process only the first N models (for quick smoke tests)",
)
parser.add_argument(
"--resume",
action="store_true",
help="Resume from previous partial run (uses progress.json)",
)
parser.add_argument(
"--no-enrich",
action="store_true",
help="Skip enrichment phase (evaluate existing metadata only)",
)
parser.add_argument(
"--audit-only",
action="store_true",
help="Run preprocessing audit only (no enrichment, no evaluation)",
)
parser.add_argument(
"--timeout",
type=int,
default=240,
help="Per-model LLM timeout in seconds (default: 240)",
)
parser.add_argument(
"-v", "--verbose",
action="store_true",
help="Enable debug logging",
)
return parser.parse_args(argv)
# ---------------------------------------------------------------------------
# Phase helpers
# ---------------------------------------------------------------------------
def _phase_header(label: str) -> None:
sep = "=" * 60
print(f"\n{sep}", file=sys.stderr)
print(f" PHASE: {label}", file=sys.stderr)
print(sep, file=sys.stderr)
# ---------------------------------------------------------------------------
# Read back LLM config after enrichment (for consistency reporting)
# ---------------------------------------------------------------------------
def _get_actual_llm_config() -> Dict[str, str]:
"""Read what LLMService is actually using, if initialized.
Only meaningful when called AFTER enrichment has started (i.e. after
``AgentService.get_instance()`` has been called).
"""
try:
from py.services.llm_service import LLMService
instance = LLMService._instance
if instance is None:
return {"status": "not initialized"}
cfg = instance._get_config()
return {
"provider": cfg.get("provider", ""),
"model": cfg.get("model", ""),
"api_base": cfg.get("api_base", ""),
}
except Exception as exc:
return {"status": f"error: {exc}"}
def _compare_llm_config(
pipeline_cfg: Dict[str, Any],
actual_cfg: Dict[str, str],
) -> List[str]:
"""Compare pipeline-loaded vs LLMService-used config.
Returns warning messages if they differ.
"""
warnings: List[str] = []
if not actual_cfg or actual_cfg.get("status", "") == "not initialized":
warnings.append(
"LLMService was not initialized during this run — cannot verify "
"config consistency."
)
return warnings
field_map = [
("llm_provider", "provider"),
("llm_model", "model"),
("llm_api_base", "api_base"),
]
for pipeline_key, llm_key in field_map:
pv = (pipeline_cfg.get(pipeline_key) or "").strip()
lv = (actual_cfg.get(llm_key) or "").strip()
if pv and lv and pv != lv:
warnings.append(
f"LLM config mismatch: --settings has '{pv}' for {pipeline_key}, "
f"but LLMService uses '{lv}'. "
f"The pipeline's --settings path ({pipeline_cfg.get('settings_path', '?')}) "
"may differ from where SettingsManager reads."
)
if not warnings and actual_cfg:
warnings.append(
"✅ LLM config matches between pipeline --settings and LLMService."
)
return warnings
# ---------------------------------------------------------------------------
# Phase 1.5: preprocessing audit
# ---------------------------------------------------------------------------
async def _run_preprocessing_audit(
entries: List[RepoEntry],
output_dir: str,
) -> Dict[str, Any]:
"""Execute the preprocessing audit and save results."""
_phase_header("Preprocessing audit")
print(f" Auditing {len(entries)} repos ...", file=sys.stderr)
readmes_dir = os.path.join(output_dir, "readmes")
t0 = time.perf_counter()
records, summary = await run_audit(entries, readmes_dir=readmes_dir)
elapsed = time.perf_counter() - t0
# Save audit data
audit_path = os.path.join(output_dir, "preprocessing_audit.json")
with open(audit_path, "w", encoding="utf-8") as fh:
json.dump(
{
"summary": summary,
"records": audit_records_to_serializable(records),
},
fh,
indent=2,
ensure_ascii=False,
)
print(f" Audit complete: {len(records)} repos in {elapsed:.0f}s", file=sys.stderr)
print(f" Section extraction activated: {summary.get('section_extraction_pct', 0)}%", file=sys.stderr)
print(f" Basename in extracted section: {summary.get('basename_in_section_pct', 0)}%", file=sys.stderr)
print(f" Avg compression: {summary.get('avg_compression_pct', 0)}%", file=sys.stderr)
print(f" Avg cleaned length: {summary.get('avg_cleaned_length', 0)} chars", file=sys.stderr)
print(f" Audit data: {audit_path}", file=sys.stderr)
if summary.get("top_flags"):
print(" Top flags:", file=sys.stderr)
for flag, count in summary["top_flags"][:5]:
print(f" - {flag}: {count}x", file=sys.stderr)
return summary
async def _run_enrichment(
model_paths: List[str],
repos: List[str],
output_dir: str,
timeout: int,
verbose: bool,
) -> Dict[str, Any]:
"""Execute the enrichment phase."""
runner = EnrichmentRunner(
output_dir=output_dir,
per_model_timeout=timeout,
)
result = await runner.run(model_paths, repos)
# Print quick summary
progress = result["progress"]
total_done = (
len(progress.get("completed", []))
+ len(progress.get("failed", []))
+ len(progress.get("timed_out", []))
)
print(
f"\n Enrichment complete: {total_done} processed "
f"({len(progress.get('completed', []))} ok, "
f"{len(progress.get('failed', []))} failed, "
f"{len(progress.get('timed_out', []))} timed out)",
file=sys.stderr,
)
return result
def _collect_enriched_metadata(
model_paths: List[str],
repos: List[str],
results: List[Dict[str, Any]],
) -> List[Dict[str, Any]]:
"""Read enriched .metadata.json for each model.
Uses the same path convention as the rest of the codebase:
``os.path.splitext(model_path)[0] + '.metadata.json'``.
Returns a list of dicts with keys: repo_id, model_path, success,
errors, metadata.
"""
enriched: List[Dict[str, Any]] = []
# Build a lookup from repo_id to enrichment result
result_lookup: Dict[str, Dict[str, Any]] = {}
for r in results:
result_lookup[r["repo_id"]] = r
for model_path, repo_id in zip(model_paths, repos):
res = result_lookup.get(repo_id, {})
metadata_path = f"{os.path.splitext(model_path)[0]}.metadata.json"
metadata: Dict[str, Any] = {}
if os.path.exists(metadata_path):
try:
with open(metadata_path, "r", encoding="utf-8") as fh:
metadata = json.load(fh)
except (json.JSONDecodeError, OSError) as exc:
logger.warning("Failed to read %s: %s", metadata_path, exc)
else:
logger.warning(
"Metadata file not found for %s (expected: %s)",
repo_id, metadata_path,
)
enriched.append({
"repo_id": repo_id,
"model_path": model_path,
"success": res.get("success", False),
"errors": res.get("errors", []),
"metadata": metadata,
})
return enriched
# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
async def main(argv: List[str]) -> int:
args = _parse_args(argv)
_setup_logging(args.verbose)
output_dir = os.path.abspath(os.path.expanduser(args.output))
os.makedirs(output_dir, exist_ok=True)
# ---- Phase 0: Initialise shared state ----
_phase_header("Initialise")
settings = load_settings(args.settings)
logger.info(
"LLM config from --settings: provider=%s model=%s api_base=%s",
settings["llm_provider"],
settings["llm_model"],
settings["llm_api_base"],
)
# Load the production base model list (replaces the old hardcoded list)
await init_supported_base_models()
# ---- Load entries ----
_phase_header("Load repo entries & construct initial metadata")
entries = load_repo_ids(args.models, max_models=args.sample if args.sample > 0 else None)
model_paths, repo_ids = create_all_initial_metadata(
entries, output_dir, skip_existing=True,
)
print(f" {len(model_paths)} repos ready", file=sys.stderr)
# ---- Phase 1.5: Preprocessing audit ----
audit_summary: Dict[str, Any] = {}
t_start = time.perf_counter()
audit_summary = await _run_preprocessing_audit(entries, output_dir)
if args.audit_only:
total_wall = time.perf_counter() - t_start
print(f"\n Audit-only done in {total_wall:.0f}s", file=sys.stderr)
print(f" Audit data: {output_dir}/preprocessing_audit.json", file=sys.stderr)
return 0
# ---- Phase 2: Enrichment ----
enrichment_results: List[Dict[str, Any]] = []
if not args.no_enrich:
_phase_header("Enrich metadata via LLM")
enrichment_out = await _run_enrichment(
model_paths, repo_ids, output_dir, args.timeout, args.verbose,
)
enrichment_results = enrichment_out["results"]
else:
print(" Enrichment skipped (--no-enrich)", file=sys.stderr)
t_enrich = time.perf_counter()
# ---- Phase 3: Evaluation ----
_phase_header("Evaluate enriched metadata")
enriched = _collect_enriched_metadata(model_paths, repo_ids, enrichment_results)
scores = evaluate_batch(enriched)
agg = aggregate_scores(scores)
print(
f" Mean total score: {agg.get('total_score', {}).get('mean', 'N/A')} / 100",
file=sys.stderr,
)
print(
f" Models scored: {agg.get('model_count', 0)}",
file=sys.stderr,
)
# ---- Phase 4: Report generation ----
_phase_header("Generate reports")
duration_summary: Dict[str, Any] | None = None
if enrichment_results:
durations = [r.get("duration_s", 0) for r in enrichment_results if r.get("duration_s")]
if durations:
sorted_d = sorted(durations)
m = len(sorted_d) // 2
duration_summary = {
"total_wall_s": round(t_enrich - t_start, 1),
"mean_s": round(sum(durations) / len(durations), 1),
"median_s": round(sorted_d[m] if len(sorted_d) % 2 else (sorted_d[m - 1] + sorted_d[m]) / 2, 1),
"min_s": round(min(durations), 1),
"max_s": round(max(durations), 1),
}
# Check LLM config consistency after enrichment (LLMService is now initialized)
actual_llm_cfg = _get_actual_llm_config()
config_warnings = _compare_llm_config(settings, actual_llm_cfg)
save_json_report(
agg, scores, enrichment_results, output_dir, duration_summary,
audit_summary=audit_summary, config_warnings=config_warnings,
)
generate_markdown_report(
agg, scores, output_dir, duration_summary,
audit_summary=audit_summary, config_warnings=config_warnings,
)
# ---- Final summary ----
total_wall = time.perf_counter() - t_start
print(f"\n Done in {total_wall:.0f}s ({total_wall / 60:.1f} min)", file=sys.stderr)
print(f" Reports: {output_dir}/report.md, {output_dir}/report.json", file=sys.stderr)
print(file=sys.stderr)
return 0 if agg.get("success_count", 0) > 0 else 1
def entry_point() -> int:
return asyncio.run(main(sys.argv[1:]))
if __name__ == "__main__":
sys.exit(entry_point())
@@ -0,0 +1,376 @@
{
"description": "Ground truth base_model mapping for HF LoRA enrichment test data",
"generated_at": "2026-07-05T18:20:00+08:00",
"inference_method": "Manual analysis of YAML base_model field + README content + filename clues",
"canonical_list_source": "Fallback list in config.py + CivitAI production API (73 models total)",
"entries": [
{
"repo_id": "k2styles/krea-2-cobalt-sky-anime-lora",
"safetensors_name": "cobalt-sky-anime.safetensors",
"yaml_base_model_raw": "krea/Krea-2-Turbo",
"correct_base_model": "Krea 2",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "k2styles/krea-2-azure-gouache-daylight-lora",
"safetensors_name": "azure-gouache-daylight.safetensors",
"yaml_base_model_raw": "krea/Krea-2-Turbo",
"correct_base_model": "Krea 2",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "TheDivergentAI/krea2-turbo-distill-lora",
"safetensors_name": "krea2_turbo_distill_r128.safetensors",
"yaml_base_model_raw": "krea/Krea-2-Raw",
"correct_base_model": "Krea 2",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "DeverStyle/Krea2-Loras",
"safetensors_name": "n0t_f4l_000001000.safetensors",
"yaml_base_model_raw": "krea/Krea-2-Turbo",
"correct_base_model": "Krea 2",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "Komorebi1995/krea2-raw-jpaf-celpaint-lora",
"safetensors_name": "krea2_raw_jpaf_celpaint_full_v1.safetensors",
"yaml_base_model_raw": null,
"correct_base_model": "Krea 2",
"confidence": "high",
"evidence": "Filename contains 'krea2'"
},
{
"repo_id": "artificialguybr/pixelartredmond-1-5v-pixel-art-loras-for-sd-1-5",
"safetensors_name": "PixelArtRedmond15V-PixelArt-PIXARFK.safetensors",
"yaml_base_model_raw": "runwayml/stable-diffusion-v1-5",
"correct_base_model": "SD 1.5",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "Shakker-Labs/FLUX.1-dev-LoRA-Logo-Design",
"safetensors_name": "FLUX-dev-lora-Logo-Design.safetensors",
"yaml_base_model_raw": "black-forest-labs/FLUX.1-dev",
"correct_base_model": "Flux.1 D",
"confidence": "high",
"evidence": "YAML base_model: FLUX.1-dev → dev → D"
},
{
"repo_id": "glif-loradex-trainer/bingbangboom_flux_surf",
"safetensors_name": "flux_surf_000001500.safetensors",
"yaml_base_model_raw": "black-forest-labs/FLUX.1-dev",
"correct_base_model": "Flux.1 D",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "prithivMLmods/Ton618-Epic-Realism-Flux-LoRA",
"safetensors_name": "Epic-Realism-Unpruned.safetensors",
"yaml_base_model_raw": "black-forest-labs/FLUX.1-dev",
"correct_base_model": "Flux.1 D",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "prithivMLmods/Fashion-Hut-Modeling-LoRA",
"safetensors_name": "Fashion-Modeling.safetensors",
"yaml_base_model_raw": "black-forest-labs/FLUX.1-dev",
"correct_base_model": "Flux.1 D",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "prithivMLmods/Retro-Pixel-Flux-LoRA",
"safetensors_name": "Retro-Pixel.safetensors",
"yaml_base_model_raw": "black-forest-labs/FLUX.1-dev",
"correct_base_model": "Flux.1 D",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "D1-3105/HiDream-E1-Full_lora",
"safetensors_name": "HiDream-E1-Full.safetensors",
"yaml_base_model_raw": "HiDream-ai/HiDream-E1-Full",
"correct_base_model": "HiDream",
"confidence": "high",
"evidence": "YAML frontmatter base_model field; filename contains 'HiDream'"
},
{
"repo_id": "renderartist/Classic-Painting-Z-Image-Turbo-LoRA",
"safetensors_name": "Classic_Painting_Z_Image_Turbo_v1_renderartist_1750.safetensors",
"yaml_base_model_raw": "Tongyi-MAI/Z-Image-Turbo",
"correct_base_model": "ZImageTurbo",
"confidence": "high",
"evidence": "YAML frontmatter base_model field; filename contains 'Z-Image-Turbo'"
},
{
"repo_id": "DeverStyle/Z-Image-loras",
"safetensors_name": "z_image_archer_style.safetensors",
"yaml_base_model_raw": "Tongyi-MAI/Z-Image-Turbo",
"correct_base_model": "ZImageTurbo",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "deadman44/Z-Image_LoRA",
"safetensors_name": "lora_zimage_turbo_myjs_alpha01.safetensors",
"yaml_base_model_raw": null,
"correct_base_model": "ZImageTurbo",
"confidence": "high",
"evidence": "Filename contains 'zimage_turbo'"
},
{
"repo_id": "zyuzuguldu/vton-lora-linen",
"safetensors_name": "pytorch_lora_weights.safetensors",
"yaml_base_model_raw": "stabilityai/stable-diffusion-xl-base-1.0",
"correct_base_model": "SDXL 1.0",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "svntax-dev/pixel_spritesheet_4walk_small_lora_v1",
"safetensors_name": "pixel_4walk_small_flux2_klein_base_4b_v1_000002750.safetensors",
"yaml_base_model_raw": "black-forest-labs/FLUX.2-klein-base-4B",
"correct_base_model": "Flux.2 Klein 4B-base",
"confidence": "high",
"evidence": "YAML base_model: FLUX.2-klein-base-4B"
},
{
"repo_id": "Haruka041/z-image-anime-lora",
"safetensors_name": "sk_anime_style_v1.0.safetensors",
"yaml_base_model_raw": "Tongyi-MAI/Z-Image-Turbo",
"correct_base_model": "ZImageTurbo",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "systms/SYSTMS-INFL8-LoRA-Wan22",
"safetensors_name": "SYSTMS_INFL8_LORA_WAN22_low_noise.safetensors",
"yaml_base_model_raw": "Wan-AI/Wan2.2-I2V-A14B",
"correct_base_model": "Wan Video 2.2 I2V-A14B",
"confidence": "high",
"evidence": "YAML base_model: Wan2.2-I2V-A14B"
},
{
"repo_id": "crafiq/flux-2-klein-9b-360-panorama-lora",
"safetensors_name": "flux-2-klein-9b-360-panorama-lora.safetensors",
"yaml_base_model_raw": "black-forest-labs/FLUX.2-klein-base-9B",
"correct_base_model": "Flux.2 Klein 9B-base",
"confidence": "high",
"evidence": "YAML base_model: FLUX.2-klein-base-9B; filename contains 'flux-2-klein-9b'"
},
{
"repo_id": "Leon1000/pixel_spritesheet_4walk_small_lora_v1",
"safetensors_name": "pixel_4walk_small_flux2_klein_base_4b_v1_000002750.safetensors",
"yaml_base_model_raw": "black-forest-labs/FLUX.2-klein-base-4B",
"correct_base_model": "Flux.2 Klein 4B-base",
"confidence": "high",
"evidence": "YAML base_model: FLUX.2-klein-base-4B; filename contains 'flux2_klein_base_4b'"
},
{
"repo_id": "Muapi/pov-missionary-legs-together-lora",
"safetensors_name": "pov-missionary-legs-together-lora.safetensors",
"yaml_base_model_raw": "OnomaAIResearch/Illustrious-xl-early-release-v0",
"correct_base_model": "Illustrious",
"confidence": "high",
"evidence": "YAML base_model: OnomaAIResearch/Illustrious-*"
},
{
"repo_id": "ostris/ideogram_4_unconditional_lora",
"safetensors_name": "ideogram_4_unconditional_lora_r16.safetensors",
"yaml_base_model_raw": "ideogram-ai/ideogram-4-fp8",
"correct_base_model": "Ideogram 4.0",
"confidence": "high",
"evidence": "YAML base_model: ideogram-ai/ideogram-4 → Ideogram 4.0; filename contains 'ideogram_4'"
},
{
"repo_id": "ilkerzgi/krea-2-bleached-surreal-uncanny-lora",
"safetensors_name": "bleached-surreal-uncanny-comfy.safetensors",
"yaml_base_model_raw": "krea/Krea-2-Turbo",
"correct_base_model": "Krea 2",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "ilkerzgi/krea-2-azure-surreal-collage-lora",
"safetensors_name": "azure-surreal-collage-comfy.safetensors",
"yaml_base_model_raw": "krea/Krea-2-Turbo",
"correct_base_model": "Krea 2",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "ilkerzgi/krea-2-airy-gouache-minimalist-lora",
"safetensors_name": "airy-gouache-minimalist-comfy.safetensors",
"yaml_base_model_raw": "krea/Krea-2-Turbo",
"correct_base_model": "Krea 2",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "k2styles/krea-2-airy-watercolor-chibi-lora",
"safetensors_name": "airy-watercolor-chibi.safetensors",
"yaml_base_model_raw": "krea/Krea-2-Turbo",
"correct_base_model": "Krea 2",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "TakeAswing/sdxl-lora-lofi",
"safetensors_name": "pytorch_lora_weights.safetensors",
"yaml_base_model_raw": "stabilityai/stable-diffusion-xl-base-1.0",
"correct_base_model": "SDXL 1.0",
"confidence": "high",
"evidence": "YAML frontmatter base_model field; repo name contains 'sdxl'"
},
{
"repo_id": "heville/anna-lora-krea2",
"safetensors_name": "pytorch_lora_weights.safetensors",
"yaml_base_model_raw": "krea/Krea-2-Raw",
"correct_base_model": "Krea 2",
"confidence": "high",
"evidence": "YAML frontmatter base_model field; repo name contains 'krea2'"
},
{
"repo_id": "Brioch/krea2_loras",
"safetensors_name": "mashap_ohwx_woman_krea2.safetensors",
"yaml_base_model_raw": "krea/Krea-2-Raw",
"correct_base_model": "Krea 2",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "hr16/Miwano-Rag-LoRA",
"safetensors_name": "Miwano-Rag-epoch10.lora.safetensors",
"yaml_base_model_raw": null,
"correct_base_model": "SD 1.5",
"confidence": "high",
"evidence": "README: base model is Kanianime (SD 1.5 fine-tune)"
},
{
"repo_id": "ikuseiso/Personal_Lora_collections",
"safetensors_name": "vergil_devil_may_cry.safetensors",
"yaml_base_model_raw": null,
"correct_base_model": "SD 1.5",
"confidence": "high",
"evidence": "Sample prompt shows Model: AbyssOrangeMix (SD 1.5), 512x768"
},
{
"repo_id": "Tanger/LoraByTanger",
"safetensors_name": "(v4)layila-000005.safetensors",
"yaml_base_model_raw": null,
"correct_base_model": "SD 1.5",
"confidence": "high",
"evidence": "README: trained on anything4.5 (SD 1.5) and nai (SD 1.5); test images on AbyssOrangeMix2_hard"
},
{
"repo_id": "DS-Archive/ds-LoRA",
"safetensors_name": "dsharu-v2_lc.safetensors",
"yaml_base_model_raw": null,
"correct_base_model": "SD 1.5",
"confidence": "high",
"evidence": "README explicitly states 'Stable Diffusion 1.5'"
},
{
"repo_id": "soknife/loras",
"safetensors_name": "irys-regular-subject-more.safetensors",
"yaml_base_model_raw": null,
"correct_base_model": "SD 1.5",
"confidence": "high",
"evidence": "README mentions SD 1.5 fine-tune models (PastelMix, AbyssOrangeMix, Anything)"
},
{
"repo_id": "prompthero/openjourney-lora",
"safetensors_name": "openjourneyLora.safetensors",
"yaml_base_model_raw": "stabilityai/stable-diffusion-2-1-base",
"correct_base_model": "SD 2.1",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "Banano/banchan-lora",
"safetensors_name": "Bananochan-PonySDXL-v2.safetensors",
"yaml_base_model_raw": null,
"correct_base_model": "Pony",
"confidence": "medium",
"evidence": "Filename contains 'PonySDXL-v2' → Pony base model"
},
{
"repo_id": "Maisman/No-Game-NoLife-LoRAs",
"safetensors_name": "ShiroNGNL2_Lora.safetensors",
"yaml_base_model_raw": null,
"correct_base_model": "SD 1.5",
"confidence": "high",
"evidence": "Sample prompts show Model: abyssorangemix2_Hardcore (SD 1.5), 512x768"
},
{
"repo_id": "EarthnDusk/Gambit_Xmen_Anime_Lora_V1.1",
"safetensors_name": "RemyLebeau.safetensors",
"yaml_base_model_raw": null,
"correct_base_model": "SD 1.5",
"confidence": "high",
"evidence": "Trained Feb 2023 via Kohya LoRA (pre-SDXL era), SD 1.5 lineage"
},
{
"repo_id": "EarthnDusk/DuskfallArt_LoRa",
"safetensors_name": "DuskfallArt.safetensors",
"yaml_base_model_raw": "stable-diffusion-v1-5/stable-diffusion-v1-5",
"correct_base_model": "SD 1.5",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "gaoxiao/pokemon-lora",
"safetensors_name": "pytorch_lora_weights.safetensors",
"yaml_base_model_raw": "runwayml/stable-diffusion-v1-5",
"correct_base_model": "SD 1.5",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "wtcherr/sd-unsplash_10k_canny-model-control-lora",
"safetensors_name": "diffusion_pytorch_model.safetensors",
"yaml_base_model_raw": "runwayml/stable-diffusion-v1-5",
"correct_base_model": "SD 1.5",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "wtcherr/sd-unsplash_10k_blur_rand_KS-model-control-lora",
"safetensors_name": "diffusion_pytorch_model.safetensors",
"yaml_base_model_raw": "runwayml/stable-diffusion-v1-5",
"correct_base_model": "SD 1.5",
"confidence": "high",
"evidence": "YAML frontmatter base_model field"
},
{
"repo_id": "samurai-architects/lora-starbucks",
"safetensors_name": "starbucks_interior.safetensors",
"yaml_base_model_raw": null,
"correct_base_model": null,
"confidence": "none",
"evidence": "README too minimal, no base_model in YAML, cannot determine"
},
{
"repo_id": "prithivMLmods/Flux-Long-Toon-LoRA",
"safetensors_name": "Long-Toon.safetensors",
"yaml_base_model_raw": "black-forest-labs/FLUX.1-dev",
"correct_base_model": "Flux.1 D",
"confidence": "high",
"evidence": "YAML base_model: FLUX.1-dev (dev → D)"
},
{
"repo_id": "Limbicnation/pixel-art-lora",
"safetensors_name": "pytorch_lora_weights.comfyui.safetensors",
"yaml_base_model_raw": "black-forest-labs/FLUX.2-klein-4B",
"correct_base_model": "Flux.2 Klein 4B",
"confidence": "high",
"evidence": "YAML base_model: FLUX.2-klein-4B; README explicitly states base model"
}
]
}
@@ -0,0 +1,46 @@
k2styles/krea-2-cobalt-sky-anime-lora, cobalt-sky-anime.safetensors
k2styles/krea-2-azure-gouache-daylight-lora, azure-gouache-daylight.safetensors
TheDivergentAI/krea2-turbo-distill-lora, krea2_turbo_distill_r128.safetensors
DeverStyle/Krea2-Loras, n0t_f4l_000001000.safetensors
Komorebi1995/krea2-raw-jpaf-celpaint-lora, krea2_raw_jpaf_celpaint_full_v1.safetensors
artificialguybr/pixelartredmond-1-5v-pixel-art-loras-for-sd-1-5, PixelArtRedmond15V-PixelArt-PIXARFK.safetensors
Shakker-Labs/FLUX.1-dev-LoRA-Logo-Design, FLUX-dev-lora-Logo-Design.safetensors
glif-loradex-trainer/bingbangboom_flux_surf, flux_surf_000001500.safetensors
prithivMLmods/Ton618-Epic-Realism-Flux-LoRA, Epic-Realism-Unpruned.safetensors
prithivMLmods/Fashion-Hut-Modeling-LoRA, Fashion-Modeling.safetensors
prithivMLmods/Retro-Pixel-Flux-LoRA, Retro-Pixel.safetensors
D1-3105/HiDream-E1-Full_lora, HiDream-E1-Full.safetensors
renderartist/Classic-Painting-Z-Image-Turbo-LoRA, Classic_Painting_Z_Image_Turbo_v1_renderartist_1750.safetensors
DeverStyle/Z-Image-loras, z_image_archer_style.safetensors
deadman44/Z-Image_LoRA, lora_zimage_turbo_myjs_alpha01.safetensors
zyuzuguldu/vton-lora-linen, pytorch_lora_weights.safetensors
svntax-dev/pixel_spritesheet_4walk_small_lora_v1, pixel_4walk_small_flux2_klein_base_4b_v1_000002750.safetensors
Haruka041/z-image-anime-lora, sk_anime_style_v1.0.safetensors
systms/SYSTMS-INFL8-LoRA-Wan22, SYSTMS_INFL8_LORA_WAN22_low_noise.safetensors
crafiq/flux-2-klein-9b-360-panorama-lora, flux-2-klein-9b-360-panorama-lora.safetensors
Leon1000/pixel_spritesheet_4walk_small_lora_v1, pixel_4walk_small_flux2_klein_base_4b_v1_000002750.safetensors
Muapi/pov-missionary-legs-together-lora, pov-missionary-legs-together-lora.safetensors
ostris/ideogram_4_unconditional_lora, ideogram_4_unconditional_lora_r16.safetensors
ilkerzgi/krea-2-bleached-surreal-uncanny-lora, bleached-surreal-uncanny-comfy.safetensors
ilkerzgi/krea-2-azure-surreal-collage-lora, azure-surreal-collage-comfy.safetensors
ilkerzgi/krea-2-airy-gouache-minimalist-lora, airy-gouache-minimalist-comfy.safetensors
k2styles/krea-2-airy-watercolor-chibi-lora, airy-watercolor-chibi.safetensors
TakeAswing/sdxl-lora-lofi, pytorch_lora_weights.safetensors
heville/anna-lora-krea2, pytorch_lora_weights.safetensors
Brioch/krea2_loras, mashap_ohwx_woman_krea2.safetensors
hr16/Miwano-Rag-LoRA, Miwano-Rag-epoch10.lora.safetensors
ikuseiso/Personal_Lora_collections, vergil_devil_may_cry.safetensors
Tanger/LoraByTanger, (v4)layila-000005.safetensors
DS-Archive/ds-LoRA, dsharu-v2_lc.safetensors
soknife/loras, irys-regular-subject-more.safetensors
prompthero/openjourney-lora, openjourneyLora.safetensors
Banano/banchan-lora, Bananochan-PonySDXL-v2.safetensors
Maisman/No-Game-NoLife-LoRAs, ShiroNGNL2_Lora.safetensors
EarthnDusk/Gambit_Xmen_Anime_Lora_V1.1, RemyLebeau.safetensors
EarthnDusk/DuskfallArt_LoRa, DuskfallArt.safetensors
gaoxiao/pokemon-lora, pytorch_lora_weights.safetensors
wtcherr/sd-unsplash_10k_canny-model-control-lora, diffusion_pytorch_model.safetensors
wtcherr/sd-unsplash_10k_blur_rand_KS-model-control-lora, diffusion_pytorch_model.safetensors
samurai-architects/lora-starbucks, starbucks_interior.safetensors
prithivMLmods/Flux-Long-Toon-LoRA, Long-Toon.safetensors
Limbicnation/pixel-art-lora, pytorch_lora_weights.comfyui.safetensors

Some files were not shown because too many files have changed in this diff Show More